Merge pull request #1536 from apernet/feat/udphop-server

feat: server side UDP port range listening (nftables/iptables)
This commit is contained in:
Toby 2026-03-29 13:24:08 -07:00 committed by GitHub
commit 82d9935c85
No known key found for this signature in database
GPG key ID: B5690EEEBB952194
10 changed files with 634 additions and 17 deletions

View file

@ -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))
}
}
}

View file

@ -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")
})
}

View 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, ",")
}

View 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]
}

View 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")
}

View 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")
}

View file

@ -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

View file

@ -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
}

View file

@ -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)
}

View 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)
}