Merge pull request #1536 from apernet/feat/udphop-server
feat: server side UDP port range listening (nftables/iptables)
This commit is contained in:
commit
82d9935c85
10 changed files with 634 additions and 17 deletions
|
|
@ -7,13 +7,16 @@ import (
|
|||
"encoding/json"
|
||||
"errors"
|
||||
"fmt"
|
||||
"io"
|
||||
"net"
|
||||
"net/http"
|
||||
"net/http/httputil"
|
||||
"net/url"
|
||||
"os"
|
||||
"os/signal"
|
||||
"strconv"
|
||||
"strings"
|
||||
"syscall"
|
||||
"time"
|
||||
|
||||
"github.com/caddyserver/certmagic"
|
||||
|
|
@ -28,6 +31,7 @@ import (
|
|||
"github.com/spf13/viper"
|
||||
"go.uber.org/zap"
|
||||
|
||||
"github.com/apernet/hysteria/app/v2/internal/firewall"
|
||||
"github.com/apernet/hysteria/app/v2/internal/utils"
|
||||
"github.com/apernet/hysteria/core/v2/server"
|
||||
"github.com/apernet/hysteria/extras/v2/auth"
|
||||
|
|
@ -271,7 +275,7 @@ func (c *serverConfig) fillConn(hyConfig *server.Config) error {
|
|||
if listenAddr == "" {
|
||||
listenAddr = defaultListenAddr
|
||||
}
|
||||
uAddr, err := net.ResolveUDPAddr("udp", listenAddr)
|
||||
uAddr, portUnion, err := resolveServerListenAddr(listenAddr)
|
||||
if err != nil {
|
||||
return configError{Field: "listen", Err: err}
|
||||
}
|
||||
|
|
@ -279,22 +283,60 @@ func (c *serverConfig) fillConn(hyConfig *server.Config) error {
|
|||
if err != nil {
|
||||
return configError{Field: "listen", Err: err}
|
||||
}
|
||||
var packetConn net.PacketConn = conn
|
||||
var cleanup io.Closer
|
||||
if len(portUnion) > 0 {
|
||||
cleanup, err = firewall.SetupUDPPortRedirect(uAddr, portUnion)
|
||||
if err != nil {
|
||||
_ = conn.Close()
|
||||
return configError{Field: "listen", Err: err}
|
||||
}
|
||||
}
|
||||
switch strings.ToLower(c.Obfs.Type) {
|
||||
case "", "plain":
|
||||
hyConfig.Conn = conn
|
||||
hyConfig.Conn = packetConn
|
||||
hyConfig.Cleanup = cleanup
|
||||
return nil
|
||||
case "salamander":
|
||||
ob, err := obfs.NewSalamanderObfuscator([]byte(c.Obfs.Salamander.Password))
|
||||
if err != nil {
|
||||
_ = conn.Close()
|
||||
if cleanup != nil {
|
||||
_ = cleanup.Close()
|
||||
}
|
||||
return configError{Field: "obfs.salamander.password", Err: err}
|
||||
}
|
||||
hyConfig.Conn = obfs.WrapPacketConn(conn, ob)
|
||||
hyConfig.Conn = obfs.WrapPacketConn(packetConn, ob)
|
||||
hyConfig.Cleanup = cleanup
|
||||
return nil
|
||||
default:
|
||||
_ = conn.Close()
|
||||
if cleanup != nil {
|
||||
_ = cleanup.Close()
|
||||
}
|
||||
return configError{Field: "obfs.type", Err: errors.New("unsupported obfuscation type")}
|
||||
}
|
||||
}
|
||||
|
||||
func resolveServerListenAddr(listenAddr string) (*net.UDPAddr, eUtils.PortUnion, error) {
|
||||
host, portStr, err := net.SplitHostPort(listenAddr)
|
||||
if err != nil {
|
||||
uAddr, resolveErr := net.ResolveUDPAddr("udp", listenAddr)
|
||||
return uAddr, nil, resolveErr
|
||||
}
|
||||
if !strings.ContainsAny(portStr, "-,") {
|
||||
uAddr, resolveErr := net.ResolveUDPAddr("udp", listenAddr)
|
||||
return uAddr, nil, resolveErr
|
||||
}
|
||||
portUnion := eUtils.ParsePortUnion(portStr)
|
||||
if portUnion == nil {
|
||||
return nil, nil, fmt.Errorf("%s is not a valid port number or range", portStr)
|
||||
}
|
||||
firstListenAddr := net.JoinHostPort(host, strconv.Itoa(int(portUnion[0].Start)))
|
||||
uAddr, err := net.ResolveUDPAddr("udp", firstListenAddr)
|
||||
return uAddr, portUnion, err
|
||||
}
|
||||
|
||||
func (c *serverConfig) fillTLSConfig(hyConfig *server.Config) error {
|
||||
if c.TLS == nil && c.ACME == nil {
|
||||
return configError{Field: "tls", Err: errors.New("must set either tls or acme")}
|
||||
|
|
@ -991,8 +1033,28 @@ func runServer(v *viper.Viper) {
|
|||
go runCheckUpdateServer()
|
||||
}
|
||||
|
||||
if err := s.Serve(); err != nil {
|
||||
logger.Fatal("failed to serve", zap.Error(err))
|
||||
signalChan := make(chan os.Signal, 1)
|
||||
signal.Notify(signalChan, os.Interrupt, syscall.SIGTERM)
|
||||
defer signal.Stop(signalChan)
|
||||
|
||||
serveErrChan := make(chan error, 1)
|
||||
go func() {
|
||||
serveErrChan <- s.Serve()
|
||||
}()
|
||||
|
||||
select {
|
||||
case <-signalChan:
|
||||
logger.Info("received signal, shutting down gracefully")
|
||||
if err := s.Close(); err != nil {
|
||||
logger.Error("failed to shut down server cleanly", zap.Error(err))
|
||||
}
|
||||
if err := <-serveErrChan; err != nil {
|
||||
logger.Info("server stopped", zap.Error(err))
|
||||
}
|
||||
case err := <-serveErrChan:
|
||||
if err != nil {
|
||||
logger.Fatal("failed to serve", zap.Error(err))
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
|
|
|
|||
|
|
@ -5,6 +5,7 @@ import (
|
|||
"time"
|
||||
|
||||
"github.com/apernet/hysteria/core/v2/server"
|
||||
eUtils "github.com/apernet/hysteria/extras/v2/utils"
|
||||
"github.com/stretchr/testify/assert"
|
||||
|
||||
"github.com/spf13/viper"
|
||||
|
|
@ -234,3 +235,24 @@ func TestServerFillCongestionConfig(t *testing.T) {
|
|||
assert.EqualError(t, err, `invalid config: congestion.bbrProfile: unsupported BBR profile "turbo"`)
|
||||
})
|
||||
}
|
||||
|
||||
func TestResolveServerListenAddr(t *testing.T) {
|
||||
t.Run("single port", func(t *testing.T) {
|
||||
addr, ports, err := resolveServerListenAddr(":8443")
|
||||
assert.NoError(t, err)
|
||||
assert.Empty(t, ports)
|
||||
assert.Equal(t, 8443, addr.Port)
|
||||
})
|
||||
|
||||
t.Run("port range", func(t *testing.T) {
|
||||
addr, ports, err := resolveServerListenAddr("127.0.0.1:9003-9001,9008")
|
||||
assert.NoError(t, err)
|
||||
assert.Equal(t, 9001, addr.Port)
|
||||
assert.Equal(t, eUtils.PortUnion{{Start: 9001, End: 9003}, {Start: 9008, End: 9008}}, ports)
|
||||
})
|
||||
|
||||
t.Run("invalid range", func(t *testing.T) {
|
||||
_, _, err := resolveServerListenAddr("127.0.0.1:9001-")
|
||||
assert.EqualError(t, err, "9001- is not a valid port number or range")
|
||||
})
|
||||
}
|
||||
|
|
|
|||
43
app/internal/firewall/firewall.go
Normal file
43
app/internal/firewall/firewall.go
Normal file
|
|
@ -0,0 +1,43 @@
|
|||
package firewall
|
||||
|
||||
import (
|
||||
"fmt"
|
||||
"net"
|
||||
"strings"
|
||||
|
||||
eUtils "github.com/apernet/hysteria/extras/v2/utils"
|
||||
)
|
||||
|
||||
type commandRunner interface {
|
||||
LookPath(file string) (string, error)
|
||||
Run(name string, args ...string) error
|
||||
}
|
||||
|
||||
func redirectPortUnion(ports eUtils.PortUnion) eUtils.PortUnion {
|
||||
if len(ports) == 0 {
|
||||
return nil
|
||||
}
|
||||
redirects := append(eUtils.PortUnion(nil), ports...)
|
||||
if redirects[0].Start == redirects[0].End {
|
||||
redirects = redirects[1:]
|
||||
} else {
|
||||
redirects[0].Start++
|
||||
}
|
||||
return redirects
|
||||
}
|
||||
|
||||
func hashInput(addr *net.UDPAddr, ports eUtils.PortUnion) string {
|
||||
return fmt.Sprintf("%s|%d|%s", addr.IP.String(), addr.Port, formatPortUnion(ports))
|
||||
}
|
||||
|
||||
func formatPortUnion(ports eUtils.PortUnion) string {
|
||||
var parts []string
|
||||
for _, portRange := range ports {
|
||||
if portRange.Start == portRange.End {
|
||||
parts = append(parts, fmt.Sprintf("%d", portRange.Start))
|
||||
} else {
|
||||
parts = append(parts, fmt.Sprintf("%d-%d", portRange.Start, portRange.End))
|
||||
}
|
||||
}
|
||||
return strings.Join(parts, ",")
|
||||
}
|
||||
217
app/internal/firewall/firewall_linux.go
Normal file
217
app/internal/firewall/firewall_linux.go
Normal file
|
|
@ -0,0 +1,217 @@
|
|||
//go:build linux
|
||||
|
||||
package firewall
|
||||
|
||||
import (
|
||||
"crypto/sha256"
|
||||
"encoding/hex"
|
||||
"fmt"
|
||||
"net"
|
||||
"os/exec"
|
||||
"sync"
|
||||
|
||||
eUtils "github.com/apernet/hysteria/extras/v2/utils"
|
||||
)
|
||||
|
||||
type osCommandRunner struct{}
|
||||
|
||||
func (osCommandRunner) LookPath(file string) (string, error) {
|
||||
return exec.LookPath(file)
|
||||
}
|
||||
|
||||
func (osCommandRunner) Run(name string, args ...string) error {
|
||||
out, err := exec.Command(name, args...).CombinedOutput()
|
||||
if err != nil {
|
||||
return fmt.Errorf("%s %v: %w: %s", name, args, err, string(out))
|
||||
}
|
||||
return nil
|
||||
}
|
||||
|
||||
type closerFuncs struct {
|
||||
mu sync.Mutex
|
||||
fns []func()
|
||||
once sync.Once
|
||||
}
|
||||
|
||||
func (c *closerFuncs) add(fn func()) {
|
||||
c.mu.Lock()
|
||||
defer c.mu.Unlock()
|
||||
c.fns = append(c.fns, fn)
|
||||
}
|
||||
|
||||
func (c *closerFuncs) Close() error {
|
||||
c.once.Do(func() {
|
||||
c.mu.Lock()
|
||||
defer c.mu.Unlock()
|
||||
for i := len(c.fns) - 1; i >= 0; i-- {
|
||||
c.fns[i]()
|
||||
}
|
||||
c.fns = nil
|
||||
})
|
||||
return nil
|
||||
}
|
||||
|
||||
func SetupUDPPortRedirect(listenAddr *net.UDPAddr, ports eUtils.PortUnion) (*closerFuncs, error) {
|
||||
cleanup, err := setupUDPPortRedirectWithRunner(osCommandRunner{}, listenAddr, ports)
|
||||
if err != nil {
|
||||
return nil, err
|
||||
}
|
||||
return cleanup, nil
|
||||
}
|
||||
|
||||
func setupUDPPortRedirectWithRunner(r commandRunner, listenAddr *net.UDPAddr, ports eUtils.PortUnion) (*closerFuncs, error) {
|
||||
redirects := redirectPortUnion(ports)
|
||||
if len(redirects) == 0 {
|
||||
return nil, nil
|
||||
}
|
||||
if _, err := r.LookPath("nft"); err == nil {
|
||||
return setupNFTablesRedirect(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{}
|
||||
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 {
|
||||
_ = cleanup.Close()
|
||||
return nil, err
|
||||
}
|
||||
cleanup.add(func() { _ = r.Run("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 {
|
||||
_ = cleanup.Close()
|
||||
return nil, err
|
||||
}
|
||||
}
|
||||
for _, chain := range []string{"prerouting", "output"} {
|
||||
for _, portRange := range redirects {
|
||||
args := []string{"add", "rule", family, tableName, chain}
|
||||
if match := nftDestinationMatch(family, listenAddr); match != nil {
|
||||
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 {
|
||||
_ = cleanup.Close()
|
||||
return nil, err
|
||||
}
|
||||
}
|
||||
}
|
||||
}
|
||||
return cleanup, nil
|
||||
}
|
||||
|
||||
func setupIPTablesRedirect(r commandRunner, listenAddr *net.UDPAddr, ports, redirects eUtils.PortUnion) (*closerFuncs, error) {
|
||||
bins, err := iptablesBinariesForAddr(r, listenAddr)
|
||||
if err != nil {
|
||||
return nil, err
|
||||
}
|
||||
cleanup := &closerFuncs{}
|
||||
for _, bin := range bins {
|
||||
hash := shortHash(bin + "|" + hashInput(listenAddr, ports))
|
||||
chainName := "HYSTERIA-PR-" + hash
|
||||
if err := r.Run(bin, "-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)
|
||||
})
|
||||
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 {
|
||||
_ = cleanup.Close()
|
||||
return nil, err
|
||||
}
|
||||
for _, baseChain := range []string{"PREROUTING", "OUTPUT"} {
|
||||
for _, portRange := range redirects {
|
||||
args := []string{"-t", "nat", "-A", baseChain}
|
||||
if match := iptablesDestinationMatch(listenAddr); match != nil {
|
||||
args = append(args, match...)
|
||||
}
|
||||
args = append(args, "-p", "udp", "--dport", iptablesPortExpr(portRange), "-j", chainName)
|
||||
if err := r.Run(bin, args...); err != nil {
|
||||
_ = cleanup.Close()
|
||||
return nil, err
|
||||
}
|
||||
deleteArgs := append([]string{"-t", "nat", "-D", baseChain}, args[4:]...)
|
||||
cleanup.add(func() { _ = r.Run(bin, deleteArgs...) })
|
||||
}
|
||||
}
|
||||
}
|
||||
return cleanup, nil
|
||||
}
|
||||
|
||||
func nftFamiliesForAddr(addr *net.UDPAddr) []string {
|
||||
if addr.IP == nil || addr.IP.IsUnspecified() {
|
||||
return []string{"ip", "ip6"}
|
||||
}
|
||||
if addr.IP.To4() != nil {
|
||||
return []string{"ip"}
|
||||
}
|
||||
return []string{"ip6"}
|
||||
}
|
||||
|
||||
func nftDestinationMatch(family string, addr *net.UDPAddr) []string {
|
||||
if addr.IP == nil || addr.IP.IsUnspecified() {
|
||||
return nil
|
||||
}
|
||||
if family == "ip6" {
|
||||
return []string{"ip6", "daddr", addr.IP.String()}
|
||||
}
|
||||
return []string{"ip", "daddr", addr.IP.String()}
|
||||
}
|
||||
|
||||
func nftPortExpr(portRange eUtils.PortRange) string {
|
||||
if portRange.Start == portRange.End {
|
||||
return fmt.Sprintf("%d", portRange.Start)
|
||||
}
|
||||
return fmt.Sprintf("%d-%d", portRange.Start, portRange.End)
|
||||
}
|
||||
|
||||
func iptablesBinariesForAddr(r commandRunner, addr *net.UDPAddr) ([]string, error) {
|
||||
if addr.IP == nil || addr.IP.IsUnspecified() {
|
||||
if _, err := r.LookPath("iptables"); err != nil {
|
||||
return nil, err
|
||||
}
|
||||
if _, err := r.LookPath("ip6tables"); err != nil {
|
||||
return nil, err
|
||||
}
|
||||
return []string{"iptables", "ip6tables"}, nil
|
||||
}
|
||||
if addr.IP.To4() != nil {
|
||||
if _, err := r.LookPath("iptables"); err != nil {
|
||||
return nil, err
|
||||
}
|
||||
return []string{"iptables"}, nil
|
||||
}
|
||||
if _, err := r.LookPath("ip6tables"); err != nil {
|
||||
return nil, err
|
||||
}
|
||||
return []string{"ip6tables"}, nil
|
||||
}
|
||||
|
||||
func iptablesDestinationMatch(addr *net.UDPAddr) []string {
|
||||
if addr.IP == nil || addr.IP.IsUnspecified() {
|
||||
return nil
|
||||
}
|
||||
return []string{"-d", addr.IP.String()}
|
||||
}
|
||||
|
||||
func iptablesPortExpr(portRange eUtils.PortRange) string {
|
||||
if portRange.Start == portRange.End {
|
||||
return fmt.Sprintf("%d", portRange.Start)
|
||||
}
|
||||
return fmt.Sprintf("%d:%d", portRange.Start, portRange.End)
|
||||
}
|
||||
|
||||
func shortHash(input string) string {
|
||||
sum := sha256.Sum256([]byte(input))
|
||||
return hex.EncodeToString(sum[:])[:8]
|
||||
}
|
||||
89
app/internal/firewall/firewall_linux_test.go
Normal file
89
app/internal/firewall/firewall_linux_test.go
Normal file
|
|
@ -0,0 +1,89 @@
|
|||
//go:build linux
|
||||
|
||||
package firewall
|
||||
|
||||
import (
|
||||
"errors"
|
||||
"net"
|
||||
"testing"
|
||||
|
||||
eUtils "github.com/apernet/hysteria/extras/v2/utils"
|
||||
"github.com/stretchr/testify/require"
|
||||
)
|
||||
|
||||
type fakeRunner struct {
|
||||
paths map[string]bool
|
||||
cmds [][]string
|
||||
fail int
|
||||
}
|
||||
|
||||
func (r *fakeRunner) LookPath(file string) (string, error) {
|
||||
if r.paths[file] {
|
||||
return "/usr/sbin/" + file, nil
|
||||
}
|
||||
return "", errors.New("not found")
|
||||
}
|
||||
|
||||
func (r *fakeRunner) Run(name string, args ...string) error {
|
||||
r.cmds = append(r.cmds, append([]string{name}, args...))
|
||||
if r.fail > 0 && len(r.cmds) == r.fail {
|
||||
return errors.New("boom")
|
||||
}
|
||||
return nil
|
||||
}
|
||||
|
||||
func TestSetupUDPPortRedirectWithRunnerNFTables(t *testing.T) {
|
||||
runner := &fakeRunner{paths: map[string]bool{"nft": true}}
|
||||
addr := &net.UDPAddr{IP: net.IPv4(1, 2, 3, 4), Port: 20000}
|
||||
ports := eUtils.PortUnion{{20000, 20002}}
|
||||
|
||||
cleanup, err := setupUDPPortRedirectWithRunner(runner, addr, ports)
|
||||
require.NoError(t, err)
|
||||
require.NotNil(t, cleanup)
|
||||
require.Contains(t, runner.cmds[0], "add")
|
||||
require.Contains(t, runner.cmds[0], "table")
|
||||
require.Contains(t, runner.cmds[3], "udp")
|
||||
require.Contains(t, runner.cmds[3], "dport")
|
||||
require.Contains(t, runner.cmds[3], "20001-20002")
|
||||
require.Contains(t, runner.cmds[3], ":20000")
|
||||
|
||||
require.NoError(t, cleanup.Close())
|
||||
require.Contains(t, runner.cmds[len(runner.cmds)-1], "delete")
|
||||
}
|
||||
|
||||
func TestSetupUDPPortRedirectWithRunnerIPTablesFallback(t *testing.T) {
|
||||
runner := &fakeRunner{paths: map[string]bool{"iptables": true}}
|
||||
addr := &net.UDPAddr{IP: net.IPv4(1, 2, 3, 4), Port: 20000}
|
||||
ports := eUtils.PortUnion{{20000, 20000}, {20002, 20003}}
|
||||
|
||||
cleanup, err := setupUDPPortRedirectWithRunner(runner, addr, ports)
|
||||
require.NoError(t, err)
|
||||
require.NotNil(t, cleanup)
|
||||
require.Equal(t, "iptables", runner.cmds[0][0])
|
||||
require.Contains(t, runner.cmds[0], "-N")
|
||||
require.Contains(t, runner.cmds[1], "REDIRECT")
|
||||
require.Contains(t, runner.cmds[2], "PREROUTING")
|
||||
require.Contains(t, runner.cmds[2], "20002:20003")
|
||||
|
||||
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
|
||||
}
|
||||
}
|
||||
require.True(t, foundDelete)
|
||||
}
|
||||
|
||||
func TestSetupUDPPortRedirectWithRunnerRollback(t *testing.T) {
|
||||
runner := &fakeRunner{
|
||||
paths: map[string]bool{"iptables": true},
|
||||
fail: 3,
|
||||
}
|
||||
addr := &net.UDPAddr{IP: net.IPv4(1, 2, 3, 4), Port: 20000}
|
||||
ports := eUtils.PortUnion{{20000, 20001}}
|
||||
|
||||
_, err := setupUDPPortRedirectWithRunner(runner, addr, ports)
|
||||
require.Error(t, err)
|
||||
require.Contains(t, runner.cmds[len(runner.cmds)-1], "-X")
|
||||
}
|
||||
15
app/internal/firewall/firewall_others.go
Normal file
15
app/internal/firewall/firewall_others.go
Normal file
|
|
@ -0,0 +1,15 @@
|
|||
//go:build !linux
|
||||
|
||||
package firewall
|
||||
|
||||
import (
|
||||
"errors"
|
||||
"io"
|
||||
"net"
|
||||
|
||||
eUtils "github.com/apernet/hysteria/extras/v2/utils"
|
||||
)
|
||||
|
||||
func SetupUDPPortRedirect(listenAddr *net.UDPAddr, ports eUtils.PortUnion) (io.Closer, error) {
|
||||
return nil, errors.New("server port-range listening is only supported on Linux")
|
||||
}
|
||||
|
|
@ -3,6 +3,7 @@ package server
|
|||
import (
|
||||
"crypto/tls"
|
||||
"crypto/x509"
|
||||
"io"
|
||||
"net"
|
||||
"net/http"
|
||||
"sync/atomic"
|
||||
|
|
@ -27,6 +28,7 @@ type Config struct {
|
|||
TLSConfig TLSConfig
|
||||
QUICConfig QUICConfig
|
||||
Conn net.PacketConn
|
||||
Cleanup io.Closer
|
||||
RequestHook RequestHook
|
||||
Outbound Outbound
|
||||
CongestionConfig CongestionConfig
|
||||
|
|
|
|||
|
|
@ -3,6 +3,7 @@ package server
|
|||
import (
|
||||
"context"
|
||||
"crypto/tls"
|
||||
"errors"
|
||||
"math/rand"
|
||||
"net/http"
|
||||
"sync"
|
||||
|
|
@ -61,7 +62,10 @@ func NewServer(config *Config) (Server, error) {
|
|||
}
|
||||
listener, err := quic.Listen(config.Conn, tlsConfig, quicConfig)
|
||||
if err != nil {
|
||||
_ = config.Conn.Close()
|
||||
err = errors.Join(err, config.Conn.Close())
|
||||
if config.Cleanup != nil {
|
||||
err = errors.Join(err, config.Cleanup.Close())
|
||||
}
|
||||
return nil, err
|
||||
}
|
||||
return &serverImpl{
|
||||
|
|
@ -86,8 +90,10 @@ func (s *serverImpl) Serve() error {
|
|||
}
|
||||
|
||||
func (s *serverImpl) Close() error {
|
||||
err := s.listener.Close()
|
||||
_ = s.config.Conn.Close()
|
||||
err := errors.Join(s.listener.Close(), s.config.Conn.Close())
|
||||
if s.config.Cleanup != nil {
|
||||
err = errors.Join(err, s.config.Cleanup.Close())
|
||||
}
|
||||
return err
|
||||
}
|
||||
|
||||
|
|
|
|||
|
|
@ -29,6 +29,9 @@ type udpHopPacketConn struct {
|
|||
|
||||
readBufferSize int
|
||||
writeBufferSize int
|
||||
deadline time.Time
|
||||
readDeadline time.Time
|
||||
writeDeadline time.Time
|
||||
|
||||
recvQueue chan *udpPacket
|
||||
closeChan chan struct{}
|
||||
|
|
@ -94,10 +97,10 @@ func (u *udpHopPacketConn) recvLoop(conn net.PacketConn) {
|
|||
u.bufPool.Put(buf)
|
||||
var netErr net.Error
|
||||
if errors.As(err, &netErr) && netErr.Timeout() {
|
||||
// Only pass through timeout errors here, not permanent errors
|
||||
// like connection closed. Connection close is normal as we close
|
||||
// the old connection to exit this loop every time we hop.
|
||||
// Pass through timeout errors, but not permanent errors such as connection closed.
|
||||
// Connection close is normal as we close the old connection to exit this loop every time we hop.
|
||||
u.recvQueue <- &udpPacket{nil, 0, nil, netErr}
|
||||
continue
|
||||
}
|
||||
return
|
||||
}
|
||||
|
|
@ -155,6 +158,15 @@ func (u *udpHopPacketConn) hop() {
|
|||
if u.writeBufferSize > 0 {
|
||||
_ = trySetWriteBuffer(u.currentConn, u.writeBufferSize)
|
||||
}
|
||||
if !u.deadline.IsZero() {
|
||||
_ = u.currentConn.SetDeadline(u.deadline)
|
||||
}
|
||||
if !u.readDeadline.IsZero() {
|
||||
_ = u.currentConn.SetReadDeadline(u.readDeadline)
|
||||
}
|
||||
if !u.writeDeadline.IsZero() {
|
||||
_ = u.currentConn.SetWriteDeadline(u.writeDeadline)
|
||||
}
|
||||
go u.recvLoop(newConn)
|
||||
// Update addrIndex to a new random value
|
||||
u.addrIndex = rand.Intn(len(u.Addrs))
|
||||
|
|
@ -215,8 +227,11 @@ func (u *udpHopPacketConn) LocalAddr() net.Addr {
|
|||
}
|
||||
|
||||
func (u *udpHopPacketConn) SetDeadline(t time.Time) error {
|
||||
u.connMutex.RLock()
|
||||
defer u.connMutex.RUnlock()
|
||||
u.connMutex.Lock()
|
||||
defer u.connMutex.Unlock()
|
||||
u.deadline = t
|
||||
u.readDeadline = t
|
||||
u.writeDeadline = t
|
||||
if u.prevConn != nil {
|
||||
_ = u.prevConn.SetDeadline(t)
|
||||
}
|
||||
|
|
@ -224,8 +239,10 @@ func (u *udpHopPacketConn) SetDeadline(t time.Time) error {
|
|||
}
|
||||
|
||||
func (u *udpHopPacketConn) SetReadDeadline(t time.Time) error {
|
||||
u.connMutex.RLock()
|
||||
defer u.connMutex.RUnlock()
|
||||
u.connMutex.Lock()
|
||||
defer u.connMutex.Unlock()
|
||||
u.deadline = time.Time{}
|
||||
u.readDeadline = t
|
||||
if u.prevConn != nil {
|
||||
_ = u.prevConn.SetReadDeadline(t)
|
||||
}
|
||||
|
|
@ -233,8 +250,10 @@ func (u *udpHopPacketConn) SetReadDeadline(t time.Time) error {
|
|||
}
|
||||
|
||||
func (u *udpHopPacketConn) SetWriteDeadline(t time.Time) error {
|
||||
u.connMutex.RLock()
|
||||
defer u.connMutex.RUnlock()
|
||||
u.connMutex.Lock()
|
||||
defer u.connMutex.Unlock()
|
||||
u.deadline = time.Time{}
|
||||
u.writeDeadline = t
|
||||
if u.prevConn != nil {
|
||||
_ = u.prevConn.SetWriteDeadline(t)
|
||||
}
|
||||
|
|
|
|||
142
extras/transport/udphop/conn_test.go
Normal file
142
extras/transport/udphop/conn_test.go
Normal file
|
|
@ -0,0 +1,142 @@
|
|||
package udphop
|
||||
|
||||
import (
|
||||
"errors"
|
||||
"net"
|
||||
"sync"
|
||||
"testing"
|
||||
"time"
|
||||
|
||||
"github.com/stretchr/testify/require"
|
||||
)
|
||||
|
||||
type timeoutError struct{}
|
||||
|
||||
func (timeoutError) Error() string { return "timeout" }
|
||||
func (timeoutError) Timeout() bool { return true }
|
||||
func (timeoutError) Temporary() bool { return true }
|
||||
|
||||
type stubPacketConn struct {
|
||||
mu sync.Mutex
|
||||
readResults []readResult
|
||||
setDeadlineCalls []time.Time
|
||||
setReadDeadlineCalls []time.Time
|
||||
setWriteDeadlineCalls []time.Time
|
||||
closed bool
|
||||
}
|
||||
|
||||
type readResult struct {
|
||||
n int
|
||||
addr net.Addr
|
||||
err error
|
||||
}
|
||||
|
||||
func (c *stubPacketConn) ReadFrom(p []byte) (int, net.Addr, error) {
|
||||
c.mu.Lock()
|
||||
defer c.mu.Unlock()
|
||||
if c.closed {
|
||||
return 0, nil, net.ErrClosed
|
||||
}
|
||||
if len(c.readResults) == 0 {
|
||||
return 0, nil, net.ErrClosed
|
||||
}
|
||||
r := c.readResults[0]
|
||||
c.readResults = c.readResults[1:]
|
||||
if r.n > 0 {
|
||||
copy(p, []byte("payload")[:r.n])
|
||||
}
|
||||
return r.n, r.addr, r.err
|
||||
}
|
||||
|
||||
func (c *stubPacketConn) WriteTo(p []byte, addr net.Addr) (int, error) { return len(p), nil }
|
||||
func (c *stubPacketConn) Close() error {
|
||||
c.mu.Lock()
|
||||
defer c.mu.Unlock()
|
||||
c.closed = true
|
||||
return nil
|
||||
}
|
||||
func (c *stubPacketConn) LocalAddr() net.Addr { return &net.UDPAddr{} }
|
||||
func (c *stubPacketConn) SetDeadline(t time.Time) error {
|
||||
c.mu.Lock()
|
||||
defer c.mu.Unlock()
|
||||
c.setDeadlineCalls = append(c.setDeadlineCalls, t)
|
||||
return nil
|
||||
}
|
||||
|
||||
func (c *stubPacketConn) SetReadDeadline(t time.Time) error {
|
||||
c.mu.Lock()
|
||||
defer c.mu.Unlock()
|
||||
c.setReadDeadlineCalls = append(c.setReadDeadlineCalls, t)
|
||||
return nil
|
||||
}
|
||||
|
||||
func (c *stubPacketConn) SetWriteDeadline(t time.Time) error {
|
||||
c.mu.Lock()
|
||||
defer c.mu.Unlock()
|
||||
c.setWriteDeadlineCalls = append(c.setWriteDeadlineCalls, t)
|
||||
return nil
|
||||
}
|
||||
|
||||
func TestRecvLoopTimeoutIsNotFatal(t *testing.T) {
|
||||
conn := &stubPacketConn{
|
||||
readResults: []readResult{
|
||||
{err: timeoutError{}},
|
||||
{n: 3, addr: &net.UDPAddr{}},
|
||||
},
|
||||
}
|
||||
u := &udpHopPacketConn{
|
||||
Addr: &net.UDPAddr{},
|
||||
recvQueue: make(chan *udpPacket, 2),
|
||||
closeChan: make(chan struct{}),
|
||||
bufPool: sync.Pool{New: func() any {
|
||||
return make([]byte, udpBufferSize)
|
||||
}},
|
||||
}
|
||||
|
||||
go u.recvLoop(conn)
|
||||
|
||||
first := <-u.recvQueue
|
||||
require.Error(t, first.Err)
|
||||
require.True(t, errors.As(first.Err, new(net.Error)))
|
||||
|
||||
second := <-u.recvQueue
|
||||
require.NoError(t, second.Err)
|
||||
require.Equal(t, 3, second.N)
|
||||
u.bufPool.Put(second.Buf)
|
||||
}
|
||||
|
||||
func TestHopReappliesStoredDeadlines(t *testing.T) {
|
||||
firstConn := &stubPacketConn{}
|
||||
secondConn := &stubPacketConn{}
|
||||
listenCalls := 0
|
||||
u := &udpHopPacketConn{
|
||||
Addr: &net.UDPAddr{},
|
||||
Addrs: []net.Addr{&net.UDPAddr{Port: 1}},
|
||||
ListenUDPFunc: func() (net.PacketConn, error) {
|
||||
listenCalls++
|
||||
if listenCalls == 1 {
|
||||
return secondConn, nil
|
||||
}
|
||||
return nil, errors.New("unexpected extra listen")
|
||||
},
|
||||
currentConn: firstConn,
|
||||
closeChan: make(chan struct{}),
|
||||
bufPool: sync.Pool{New: func() any {
|
||||
return make([]byte, udpBufferSize)
|
||||
}},
|
||||
}
|
||||
|
||||
deadline := time.Now().Add(time.Minute)
|
||||
readDeadline := time.Now().Add(2 * time.Minute)
|
||||
writeDeadline := time.Now().Add(3 * time.Minute)
|
||||
|
||||
require.NoError(t, u.SetDeadline(deadline))
|
||||
require.NoError(t, u.SetReadDeadline(readDeadline))
|
||||
require.NoError(t, u.SetWriteDeadline(writeDeadline))
|
||||
|
||||
u.hop()
|
||||
|
||||
require.Empty(t, secondConn.setDeadlineCalls)
|
||||
require.Equal(t, []time.Time{readDeadline}, secondConn.setReadDeadlineCalls)
|
||||
require.Equal(t, []time.Time{writeDeadline}, secondConn.setWriteDeadlineCalls)
|
||||
}
|
||||
Loading…
Add table
Reference in a new issue