diff --git a/app/cmd/server.go b/app/cmd/server.go index 8a05947..4e23c21 100644 --- a/app/cmd/server.go +++ b/app/cmd/server.go @@ -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)) + } } } diff --git a/app/cmd/server_test.go b/app/cmd/server_test.go index 2be6ade..a9cb38b 100644 --- a/app/cmd/server_test.go +++ b/app/cmd/server_test.go @@ -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") + }) +} diff --git a/app/internal/firewall/firewall.go b/app/internal/firewall/firewall.go new file mode 100644 index 0000000..81ac9a8 --- /dev/null +++ b/app/internal/firewall/firewall.go @@ -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, ",") +} diff --git a/app/internal/firewall/firewall_linux.go b/app/internal/firewall/firewall_linux.go new file mode 100644 index 0000000..9730b06 --- /dev/null +++ b/app/internal/firewall/firewall_linux.go @@ -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] +} diff --git a/app/internal/firewall/firewall_linux_test.go b/app/internal/firewall/firewall_linux_test.go new file mode 100644 index 0000000..12da9b9 --- /dev/null +++ b/app/internal/firewall/firewall_linux_test.go @@ -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") +} diff --git a/app/internal/firewall/firewall_others.go b/app/internal/firewall/firewall_others.go new file mode 100644 index 0000000..6910bba --- /dev/null +++ b/app/internal/firewall/firewall_others.go @@ -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") +} diff --git a/core/server/config.go b/core/server/config.go index 1dde39f..1d5c568 100644 --- a/core/server/config.go +++ b/core/server/config.go @@ -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 diff --git a/core/server/server.go b/core/server/server.go index 4821717..ed5160a 100644 --- a/core/server/server.go +++ b/core/server/server.go @@ -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 } diff --git a/extras/transport/udphop/conn.go b/extras/transport/udphop/conn.go index 32cc31c..5039ed9 100644 --- a/extras/transport/udphop/conn.go +++ b/extras/transport/udphop/conn.go @@ -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) } diff --git a/extras/transport/udphop/conn_test.go b/extras/transport/udphop/conn_test.go new file mode 100644 index 0000000..4c1333f --- /dev/null +++ b/extras/transport/udphop/conn_test.go @@ -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) +}