From 737ebf01c43c3aaf59d71624c2032674e011a004 Mon Sep 17 00:00:00 2001 From: =?UTF-8?q?=E4=B8=96=E7=95=8C?= Date: Sun, 24 Aug 2025 17:46:03 +0800 Subject: [PATCH] ping: Add timeout to destinations --- ping/destination.go | 27 +++++++++++++++++++++++---- ping/destination_gvisor.go | 16 ++++++++++++++++ ping/destination_test.go | 24 ++++++++++++++++++++++++ ping/ping.go | 6 ++++++ ping/socket_linux_unprivileged.go | 23 ++++++++--------------- route_direct.go | 4 ++++ 6 files changed, 81 insertions(+), 19 deletions(-) create mode 100644 ping/destination_test.go diff --git a/ping/destination.go b/ping/destination.go index 3b87c74..dc20112 100644 --- a/ping/destination.go +++ b/ping/destination.go @@ -6,6 +6,7 @@ import ( "net/netip" "os" "runtime" + "time" "github.com/sagernet/sing-tun" "github.com/sagernet/sing/common/buf" @@ -17,13 +18,21 @@ import ( var _ tun.DirectRouteDestination = (*Destination)(nil) type Destination struct { + conn *Conn ctx context.Context logger logger.ContextLogger routeContext tun.DirectRouteContext - conn *Conn + timeout time.Duration } -func ConnectDestination(ctx context.Context, logger logger.ContextLogger, controlFunc control.Func, address netip.Addr, routeContext tun.DirectRouteContext) (tun.DirectRouteDestination, error) { +func ConnectDestination( + ctx context.Context, + logger logger.ContextLogger, + controlFunc control.Func, + address netip.Addr, + routeContext tun.DirectRouteContext, + timeout time.Duration, +) (tun.DirectRouteDestination, error) { var ( conn *Conn err error @@ -41,19 +50,25 @@ func ConnectDestination(ctx context.Context, logger logger.ContextLogger, contro return nil, err } d := &Destination{ + conn: conn, ctx: ctx, logger: logger, routeContext: routeContext, - conn: conn, + timeout: timeout, } go d.loopRead() return d, nil } func (d *Destination) loopRead() { + defer d.Close() for { buffer := buf.NewPacket() - err := d.conn.ReadIP(buffer) + err := d.conn.SetReadDeadline(time.Now().Add(d.timeout)) + if err != nil { + d.logger.ErrorContext(d.ctx, E.Cause(err, "set read deadline for ICMP conn")) + } + err = d.conn.ReadIP(buffer) if err != nil { buffer.Release() if !E.IsClosed(err) { @@ -76,3 +91,7 @@ func (d *Destination) WritePacket(packet *buf.Buffer) error { func (d *Destination) Close() error { return d.conn.Close() } + +func (d *Destination) IsClosed() bool { + return d.conn.IsClosed() +} diff --git a/ping/destination_gvisor.go b/ping/destination_gvisor.go index abe98c8..a026f4a 100644 --- a/ping/destination_gvisor.go +++ b/ping/destination_gvisor.go @@ -5,11 +5,13 @@ package ping import ( "context" "net/netip" + "time" "github.com/sagernet/gvisor/pkg/tcpip" "github.com/sagernet/gvisor/pkg/tcpip/adapters/gonet" "github.com/sagernet/gvisor/pkg/tcpip/header" "github.com/sagernet/gvisor/pkg/tcpip/stack" + "github.com/sagernet/gvisor/pkg/tcpip/transport" "github.com/sagernet/gvisor/pkg/waiter" "github.com/sagernet/sing-tun" "github.com/sagernet/sing/common" @@ -23,8 +25,10 @@ var _ tun.DirectRouteDestination = (*GVisorDestination)(nil) type GVisorDestination struct { ctx context.Context logger logger.ContextLogger + endpoint tcpip.Endpoint conn *gonet.TCPConn rewriter *Rewriter + timeout time.Duration } func ConnectGVisor( @@ -33,6 +37,7 @@ func ConnectGVisor( routeContext tun.DirectRouteContext, stack *stack.Stack, bindAddress4, bindAddress6 netip.Addr, + timeout time.Duration, ) (*GVisorDestination, error) { var ( bindAddress tcpip.Address @@ -76,16 +81,23 @@ func ConnectGVisor( destination := &GVisorDestination{ ctx: ctx, logger: logger, + endpoint: endpoint, conn: gonet.NewTCPConn(&wq, endpoint), rewriter: rewriter, + timeout: timeout, } go destination.loopRead() return destination, nil } func (d *GVisorDestination) loopRead() { + defer d.endpoint.Close() for { buffer := buf.NewPacket() + err := d.conn.SetReadDeadline(time.Now().Add(d.timeout)) + if err != nil { + d.logger.ErrorContext(d.ctx, E.Cause(err, "set read deadline for ICMP conn")) + } n, err := d.conn.Read(buffer.FreeBytes()) if err != nil { buffer.Release() @@ -111,3 +123,7 @@ func (d *GVisorDestination) WritePacket(packet *buf.Buffer) error { func (d *GVisorDestination) Close() error { return d.conn.Close() } + +func (d *GVisorDestination) IsClosed() bool { + return transport.DatagramEndpointState(d.endpoint.State()) == transport.DatagramEndpointStateClosed +} diff --git a/ping/destination_test.go b/ping/destination_test.go new file mode 100644 index 0000000..d0a1af8 --- /dev/null +++ b/ping/destination_test.go @@ -0,0 +1,24 @@ +package ping_test + +import ( + "context" + "net/netip" + "testing" + "time" + + "github.com/sagernet/sing-tun/ping" + "github.com/sagernet/sing/common/logger" + + "github.com/stretchr/testify/require" +) + +func TestIsClosed(t *testing.T) { + t.Parallel() + destination, err := ping.ConnectDestination(context.Background(), logger.NOP(), nil, netip.MustParseAddr("1.1.1.1"), nil, 30*time.Second) + require.NoError(t, err) + defer destination.Close() + time.Sleep(1 * time.Second) + require.False(t, destination.IsClosed()) + destination.Close() + require.True(t, destination.IsClosed()) +} diff --git a/ping/ping.go b/ping/ping.go index cbe0672..1ea0edf 100644 --- a/ping/ping.go +++ b/ping/ping.go @@ -29,6 +29,7 @@ type Conn struct { conn net.Conn destination netip.Addr source atomic.TypedValue[netip.Addr] + closed atomic.Bool } func Connect(ctx context.Context, logger logger.ContextLogger, privileged bool, controlFunc control.Func, destination netip.Addr) (*Conn, error) { @@ -230,5 +231,10 @@ func (c *Conn) SetReadDeadline(t time.Time) error { } func (c *Conn) Close() error { + defer c.closed.Store(true) return c.conn.Close() } + +func (c *Conn) IsClosed() bool { + return c.closed.Load() +} diff --git a/ping/socket_linux_unprivileged.go b/ping/socket_linux_unprivileged.go index 6059347..1da0c83 100644 --- a/ping/socket_linux_unprivileged.go +++ b/ping/socket_linux_unprivileged.go @@ -16,13 +16,12 @@ import ( ) type UnprivilegedConn struct { - ctx context.Context - cancel context.CancelFunc - controlFunc control.Func - destination netip.Addr - receiveChan chan *unprivilegedResponse - readDeadline atomic.TypedValue[time.Time] - writeDeadline atomic.TypedValue[time.Time] + ctx context.Context + cancel context.CancelFunc + controlFunc control.Func + destination netip.Addr + receiveChan chan *unprivilegedResponse + readDeadline atomic.TypedValue[time.Time] } type unprivilegedResponse struct { @@ -89,9 +88,6 @@ func (c *UnprivilegedConn) Write(b []byte) (n int, err error) { if readDeadline := c.readDeadline.Load(); !readDeadline.IsZero() { conn.SetReadDeadline(readDeadline) } - if writeDeadline := c.writeDeadline.Load(); !writeDeadline.IsZero() { - conn.SetWriteDeadline(writeDeadline) - } n, err = conn.Write(b) if err != nil { conn.Close() @@ -157,9 +153,7 @@ func (c *UnprivilegedConn) RemoteAddr() net.Addr { } func (c *UnprivilegedConn) SetDeadline(t time.Time) error { - c.readDeadline.Store(t) - c.writeDeadline.Store(t) - return nil + return os.ErrInvalid } func (c *UnprivilegedConn) SetReadDeadline(t time.Time) error { @@ -168,6 +162,5 @@ func (c *UnprivilegedConn) SetReadDeadline(t time.Time) error { } func (c *UnprivilegedConn) SetWriteDeadline(t time.Time) error { - c.writeDeadline.Store(t) - return nil + return os.ErrInvalid } diff --git a/route_direct.go b/route_direct.go index 8358ae5..2279aa8 100644 --- a/route_direct.go +++ b/route_direct.go @@ -13,6 +13,7 @@ import ( type DirectRouteDestination interface { WritePacket(packet *buf.Buffer) error Close() error + IsClosed() bool } type DirectRouteSession struct { @@ -28,6 +29,9 @@ type DirectRouteMapping struct { func NewDirectRouteMapping(timeout time.Duration) *DirectRouteMapping { mapping := common.Must1(freelru.NewSharded[DirectRouteSession, DirectRouteDestination](1024, maphash.NewHasher[DirectRouteSession]().Hash32)) + mapping.SetHealthCheck(func(session DirectRouteSession, destination DirectRouteDestination) bool { + return !destination.IsClosed() + }) mapping.SetOnEvict(func(session DirectRouteSession, action DirectRouteDestination) { action.Close() })