From aae2f0750c3fc7f7f336e6ec3365c4d1e68f3396 Mon Sep 17 00:00:00 2001 From: =?UTF-8?q?=E4=B8=96=E7=95=8C?= Date: Mon, 29 Jun 2026 10:04:52 +0800 Subject: [PATCH] Fix ping --- ping/ping.go | 7 ++- ping/ping_test.go | 114 +++++++++++++++++++++++++++++++++++++++++++ stack_gvisor.go | 2 +- stack_gvisor_icmp.go | 25 +++++++--- 4 files changed, 137 insertions(+), 11 deletions(-) diff --git a/ping/ping.go b/ping/ping.go index d855424..81ac696 100644 --- a/ping/ping.go +++ b/ping/ping.go @@ -24,6 +24,7 @@ type Conn struct { ctx context.Context privileged bool conn net.Conn + controlConn net.Conn destination netip.Addr source common.TypedValue[netip.Addr] closed atomic.Bool @@ -53,6 +54,7 @@ func (c *Conn) connect(controlFunc control.Func, idleTimeout time.Duration) (err return err } if ipConn, isIPConn := common.Cast[*net.IPConn](c.conn); isIPConn { + c.controlConn = ipConn c.readMsg = func(b, oob []byte) (n, oobn int, addr netip.Addr, err error) { var ipAddr *net.IPAddr n, oobn, _, ipAddr, err = ipConn.ReadMsgIP(b, oob) @@ -62,6 +64,7 @@ func (c *Conn) connect(controlFunc control.Func, idleTimeout time.Duration) (err return } } else if udpConn, isUDPConn := common.Cast[*net.UDPConn](c.conn); isUDPConn { + c.controlConn = udpConn c.readMsg = func(b, oob []byte) (n, oobn int, addr netip.Addr, err error) { var addrPort netip.AddrPort n, oobn, _, addrPort, err = udpConn.ReadMsgUDPAddrPort(b, oob) @@ -258,7 +261,7 @@ func (c *Conn) WriteIP(buffer *buf.Buffer) error { if !c.destination.Is6() { ipHdr := header.IPv4(buffer.Bytes()) if !c.isLinuxUnprivileged() { - err := ipv4.NewConn(c.conn).SetTTL(int(ipHdr.TTL())) + err := ipv4.NewConn(c.controlConn).SetTTL(int(ipHdr.TTL())) if err != nil { return err } @@ -271,7 +274,7 @@ func (c *Conn) WriteIP(buffer *buf.Buffer) error { } else { ipHdr := header.IPv6(buffer.Bytes()) if !c.isLinuxUnprivileged() { - err := ipv6.NewConn(c.conn).SetHopLimit(int(ipHdr.HopLimit())) + err := ipv6.NewConn(c.controlConn).SetHopLimit(int(ipHdr.HopLimit())) if err != nil { return err } diff --git a/ping/ping_test.go b/ping/ping_test.go index 5a04be1..9fef5bf 100644 --- a/ping/ping_test.go +++ b/ping/ping_test.go @@ -30,6 +30,9 @@ func TestPing(t *testing.T) { t.Run("read-ip", func(t *testing.T) { testPingIPv4ReadIP(t, false, addr4) }) + t.Run("write-ip", func(t *testing.T) { + testPingIPv4WriteIP(t, false, addr4) + }) }) t.Run("privileged", func(t *testing.T) { if runtime.GOOS != "windows" && os.Getuid() != 0 { @@ -41,6 +44,9 @@ func TestPing(t *testing.T) { t.Run("read-ip", func(t *testing.T) { testPingIPv4ReadIP(t, true, addr4) }) + t.Run("write-ip", func(t *testing.T) { + testPingIPv4WriteIP(t, true, addr4) + }) }) }) // const addr6 = "2606:4700:4700::1001" @@ -56,6 +62,9 @@ func TestPing(t *testing.T) { t.Run("read-ip", func(t *testing.T) { testPingIPv6ReadIP(t, false, addr6) }) + t.Run("write-ip", func(t *testing.T) { + testPingIPv6WriteIP(t, false, addr6) + }) }) t.Run("privileged", func(t *testing.T) { if runtime.GOOS != "windows" && os.Getuid() != 0 { @@ -67,6 +76,9 @@ func TestPing(t *testing.T) { t.Run("read-ip", func(t *testing.T) { testPingIPv6ReadIP(t, true, addr6) }) + t.Run("write-ip", func(t *testing.T) { + testPingIPv6WriteIP(t, true, addr6) + }) }) }) } @@ -196,3 +208,105 @@ func testPingIPv6ReadICMP(t *testing.T, privileged bool, addr string) { require.Equal(t, header.ICMPv6EchoReply, icmpHdr.Type()) require.Equal(t, request.Ident(), icmpHdr.Ident()) } + +// testPingIPv4WriteIP exercises the WriteIP send path, which is what the real +// TUN flow uses (Destination.WritePacket -> Conn.WriteIP). Unlike the ReadIP/ +// ReadICMP tests, which send via WriteICMP, this path runs the SetTTL call that +// regressed in ebb52fb: on privileged Linux / unprivileged macOS the socket is +// wrapped in a BindPacketConn that does not expose SyscallConn, so +// ipv4.NewConn(c.conn).SetTTL returned "invalid connection" and the echo +// request was never sent. +func testPingIPv4WriteIP(t *testing.T, privileged bool, addr string) { + conn, err := ping.Connect(context.Background(), privileged, nil, netip.MustParseAddr(addr), 0) + if runtime.GOOS == "linux" && err != nil && err.Error() == "socket(): permission denied" { + t.SkipNow() + } + require.NoError(t, err) + defer conn.Close() + + ident := uint16(rand.Uint32()) + const totalLen = header.IPv4MinimumSize + header.ICMPv4MinimumSize + packet := buf.NewSize(totalLen) + ipHdr := header.IPv4(packet.Extend(totalLen)) + ipHdr.Encode(&header.IPv4Fields{ + TotalLength: totalLen, + TTL: 64, + Protocol: uint8(header.ICMPv4ProtocolNumber), + SrcAddr: netip.MustParseAddr("127.0.0.1"), + DstAddr: netip.MustParseAddr(addr), + }) + ipHdr.SetChecksum(^ipHdr.CalculateChecksum()) + icmpHdr := header.ICMPv4(ipHdr.Payload()) + icmpHdr.SetType(header.ICMPv4Echo) + icmpHdr.SetIdent(ident) + icmpHdr.SetChecksum(header.ICMPv4Checksum(icmpHdr, 0)) + + conn.SetLocalAddr(netip.MustParseAddr("127.0.0.1")) + err = conn.WriteIP(packet) + require.NoError(t, err, "WriteIP must send the echo request") + + require.NoError(t, conn.SetReadDeadline(time.Now().Add(3*time.Second))) + response := buf.NewPacket() + defer response.Release() + err = conn.ReadIP(response) + require.NoError(t, err) + if runtime.GOOS == "linux" && privileged { + response.Reset() + err = conn.ReadIP(response) + require.NoError(t, err) + } + respIP := header.IPv4(response.Bytes()) + require.NotZero(t, respIP.TTL()) + respICMP := header.ICMPv4(respIP.Payload()) + require.Equal(t, header.ICMPv4EchoReply, respICMP.Type()) + require.Equal(t, ident, respICMP.Ident()) +} + +func testPingIPv6WriteIP(t *testing.T, privileged bool, addr string) { + conn, err := ping.Connect(context.Background(), privileged, nil, netip.MustParseAddr(addr), 0) + if runtime.GOOS == "linux" && err != nil && err.Error() == "socket(): permission denied" { + t.SkipNow() + } + require.NoError(t, err) + defer conn.Close() + + ident := uint16(rand.Uint32()) + const payloadLen = header.ICMPv6MinimumSize + packet := buf.NewSize(header.IPv6MinimumSize + payloadLen) + ipHdr := header.IPv6(packet.Extend(header.IPv6MinimumSize + payloadLen)) + ipHdr.Encode(&header.IPv6Fields{ + PayloadLength: payloadLen, + TransportProtocol: header.ICMPv6ProtocolNumber, + HopLimit: 64, + SrcAddr: netip.MustParseAddr("::1"), + DstAddr: netip.MustParseAddr(addr), + }) + icmpHdr := header.ICMPv6(ipHdr.Payload()) + icmpHdr.SetType(header.ICMPv6EchoRequest) + icmpHdr.SetIdent(ident) + icmpHdr.SetChecksum(header.ICMPv6Checksum(header.ICMPv6ChecksumParams{ + Header: icmpHdr, + Src: ipHdr.SourceAddressSlice(), + Dst: ipHdr.DestinationAddressSlice(), + })) + + conn.SetLocalAddr(netip.MustParseAddr("::1")) + err = conn.WriteIP(packet) + require.NoError(t, err, "WriteIP must send the echo request") + + require.NoError(t, conn.SetReadDeadline(time.Now().Add(3*time.Second))) + response := buf.NewPacket() + defer response.Release() + err = conn.ReadIP(response) + require.NoError(t, err) + if runtime.GOOS == "darwin" || runtime.GOOS == "linux" && privileged { + response.Reset() + err = conn.ReadIP(response) + require.NoError(t, err) + } + respIP := header.IPv6(response.Bytes()) + require.NotZero(t, respIP.HopLimit()) + respICMP := header.ICMPv6(respIP.Payload()) + require.Equal(t, header.ICMPv6EchoReply, respICMP.Type()) + require.Equal(t, ident, respICMP.Ident()) +} diff --git a/stack_gvisor.go b/stack_gvisor.go index 63df41a..92e84ba 100644 --- a/stack_gvisor.go +++ b/stack_gvisor.go @@ -95,7 +95,7 @@ func (t *GVisor) Start() error { } ipStack.SetTransportProtocolHandler(tcp.ProtocolNumber, NewTCPForwarderWithLoopback(t.ctx, ipStack, t.handler, t.inet4LoopbackAddress, t.inet6LoopbackAddress, t.tun).HandlePacket) ipStack.SetTransportProtocolHandler(udp.ProtocolNumber, NewUDPForwarder(t.ctx, ipStack, t.handler, t.udpTimeout).HandlePacket) - icmpForwarder := NewICMPForwarder(t.ctx, ipStack, t.handler, t.icmpTimeout) + icmpForwarder := NewICMPForwarder(t.ctx, ipStack, t.logger, t.handler, t.icmpTimeout) icmpForwarder.SetLocalAddresses(t.inet4Address, t.inet6Address) ipStack.SetTransportProtocolHandler(icmp.ProtocolNumber4, icmpForwarder.HandlePacket) ipStack.SetTransportProtocolHandler(icmp.ProtocolNumber6, icmpForwarder.HandlePacket) diff --git a/stack_gvisor_icmp.go b/stack_gvisor_icmp.go index da5549b..1f11bbd 100644 --- a/stack_gvisor_icmp.go +++ b/stack_gvisor_icmp.go @@ -18,6 +18,8 @@ import ( "github.com/sagernet/gvisor/pkg/tcpip/network/ipv6" "github.com/sagernet/gvisor/pkg/tcpip/stack" "github.com/sagernet/sing/common/buf" + E "github.com/sagernet/sing/common/exceptions" + "github.com/sagernet/sing/common/logger" M "github.com/sagernet/sing/common/metadata" N "github.com/sagernet/sing/common/network" ) @@ -25,6 +27,7 @@ import ( type ICMPForwarder struct { ctx context.Context stack *stack.Stack + logger logger.Logger inet4Address netip.Addr inet6Address netip.Addr handler Handler @@ -34,12 +37,14 @@ type ICMPForwarder struct { func NewICMPForwarder( ctx context.Context, stack *stack.Stack, + logger logger.Logger, handler Handler, timeout time.Duration, ) *ICMPForwarder { return &ICMPForwarder{ ctx: ctx, stack: stack, + logger: logger, handler: handler, mapping: NewDirectRouteMapping(timeout), } @@ -81,8 +86,10 @@ func (f *ICMPForwarder) HandlePacket(id stack.TransportEndpointID, pkt *stack.Pa return true } if action != nil { - // TODO: handle error - _ = icmpWritePacketBuffer(action, pkt) + err = icmpWritePacketBuffer(action, pkt) + if err != nil { + f.logger.Error(E.Cause(err, "write ICMPv4 echo request")) + } return true } } @@ -95,7 +102,7 @@ func (f *ICMPForwarder) HandlePacket(id stack.TransportEndpointID, pkt *stack.Pa ipHdr.SetChecksum(^ipHdr.CalculateChecksum()) outgoingEP, gErr := f.stack.GetNetworkEndpoint(DefaultNIC, header.IPv4ProtocolNumber) if gErr != nil { - // TODO: log error + f.logger.Error(E.Cause(gonet.TranslateNetstackError(gErr), "get IPv4 network endpoint")) return true } route, gErr := f.stack.FindRoute( @@ -106,7 +113,7 @@ func (f *ICMPForwarder) HandlePacket(id stack.TransportEndpointID, pkt *stack.Pa false, ) if gErr != nil { - // TODO: log error + f.logger.Error(E.Cause(gonet.TranslateNetstackError(gErr), "find IPv4 route")) return true } defer route.Release() @@ -142,9 +149,11 @@ func (f *ICMPForwarder) HandlePacket(id stack.TransportEndpointID, pkt *stack.Pa return true } if action != nil { - // TODO: handle error pkt.IncRef() - _ = icmpWritePacketBuffer(action, pkt) + err = icmpWritePacketBuffer(action, pkt) + if err != nil { + f.logger.Error(E.Cause(err, "write ICMPv6 echo request")) + } return true } } @@ -161,7 +170,7 @@ func (f *ICMPForwarder) HandlePacket(id stack.TransportEndpointID, pkt *stack.Pa })) outgoingEP, gErr := f.stack.GetNetworkEndpoint(DefaultNIC, header.IPv4ProtocolNumber) if gErr != nil { - // TODO: log error + f.logger.Error(E.Cause(gonet.TranslateNetstackError(gErr), "get IPv6 network endpoint")) return true } route, gErr := f.stack.FindRoute( @@ -172,7 +181,7 @@ func (f *ICMPForwarder) HandlePacket(id stack.TransportEndpointID, pkt *stack.Pa false, ) if gErr != nil { - // TODO: log error + f.logger.Error(E.Cause(gonet.TranslateNetstackError(gErr), "find IPv6 route")) return true } defer route.Release()