Fix ping
This commit is contained in:
parent
1251022fce
commit
aae2f0750c
4 changed files with 137 additions and 11 deletions
|
|
@ -24,6 +24,7 @@ type Conn struct {
|
||||||
ctx context.Context
|
ctx context.Context
|
||||||
privileged bool
|
privileged bool
|
||||||
conn net.Conn
|
conn net.Conn
|
||||||
|
controlConn net.Conn
|
||||||
destination netip.Addr
|
destination netip.Addr
|
||||||
source common.TypedValue[netip.Addr]
|
source common.TypedValue[netip.Addr]
|
||||||
closed atomic.Bool
|
closed atomic.Bool
|
||||||
|
|
@ -53,6 +54,7 @@ func (c *Conn) connect(controlFunc control.Func, idleTimeout time.Duration) (err
|
||||||
return err
|
return err
|
||||||
}
|
}
|
||||||
if ipConn, isIPConn := common.Cast[*net.IPConn](c.conn); isIPConn {
|
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) {
|
c.readMsg = func(b, oob []byte) (n, oobn int, addr netip.Addr, err error) {
|
||||||
var ipAddr *net.IPAddr
|
var ipAddr *net.IPAddr
|
||||||
n, oobn, _, ipAddr, err = ipConn.ReadMsgIP(b, oob)
|
n, oobn, _, ipAddr, err = ipConn.ReadMsgIP(b, oob)
|
||||||
|
|
@ -62,6 +64,7 @@ func (c *Conn) connect(controlFunc control.Func, idleTimeout time.Duration) (err
|
||||||
return
|
return
|
||||||
}
|
}
|
||||||
} else if udpConn, isUDPConn := common.Cast[*net.UDPConn](c.conn); isUDPConn {
|
} 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) {
|
c.readMsg = func(b, oob []byte) (n, oobn int, addr netip.Addr, err error) {
|
||||||
var addrPort netip.AddrPort
|
var addrPort netip.AddrPort
|
||||||
n, oobn, _, addrPort, err = udpConn.ReadMsgUDPAddrPort(b, oob)
|
n, oobn, _, addrPort, err = udpConn.ReadMsgUDPAddrPort(b, oob)
|
||||||
|
|
@ -258,7 +261,7 @@ func (c *Conn) WriteIP(buffer *buf.Buffer) error {
|
||||||
if !c.destination.Is6() {
|
if !c.destination.Is6() {
|
||||||
ipHdr := header.IPv4(buffer.Bytes())
|
ipHdr := header.IPv4(buffer.Bytes())
|
||||||
if !c.isLinuxUnprivileged() {
|
if !c.isLinuxUnprivileged() {
|
||||||
err := ipv4.NewConn(c.conn).SetTTL(int(ipHdr.TTL()))
|
err := ipv4.NewConn(c.controlConn).SetTTL(int(ipHdr.TTL()))
|
||||||
if err != nil {
|
if err != nil {
|
||||||
return err
|
return err
|
||||||
}
|
}
|
||||||
|
|
@ -271,7 +274,7 @@ func (c *Conn) WriteIP(buffer *buf.Buffer) error {
|
||||||
} else {
|
} else {
|
||||||
ipHdr := header.IPv6(buffer.Bytes())
|
ipHdr := header.IPv6(buffer.Bytes())
|
||||||
if !c.isLinuxUnprivileged() {
|
if !c.isLinuxUnprivileged() {
|
||||||
err := ipv6.NewConn(c.conn).SetHopLimit(int(ipHdr.HopLimit()))
|
err := ipv6.NewConn(c.controlConn).SetHopLimit(int(ipHdr.HopLimit()))
|
||||||
if err != nil {
|
if err != nil {
|
||||||
return err
|
return err
|
||||||
}
|
}
|
||||||
|
|
|
||||||
|
|
@ -30,6 +30,9 @@ func TestPing(t *testing.T) {
|
||||||
t.Run("read-ip", func(t *testing.T) {
|
t.Run("read-ip", func(t *testing.T) {
|
||||||
testPingIPv4ReadIP(t, false, addr4)
|
testPingIPv4ReadIP(t, false, addr4)
|
||||||
})
|
})
|
||||||
|
t.Run("write-ip", func(t *testing.T) {
|
||||||
|
testPingIPv4WriteIP(t, false, addr4)
|
||||||
|
})
|
||||||
})
|
})
|
||||||
t.Run("privileged", func(t *testing.T) {
|
t.Run("privileged", func(t *testing.T) {
|
||||||
if runtime.GOOS != "windows" && os.Getuid() != 0 {
|
if runtime.GOOS != "windows" && os.Getuid() != 0 {
|
||||||
|
|
@ -41,6 +44,9 @@ func TestPing(t *testing.T) {
|
||||||
t.Run("read-ip", func(t *testing.T) {
|
t.Run("read-ip", func(t *testing.T) {
|
||||||
testPingIPv4ReadIP(t, true, addr4)
|
testPingIPv4ReadIP(t, true, addr4)
|
||||||
})
|
})
|
||||||
|
t.Run("write-ip", func(t *testing.T) {
|
||||||
|
testPingIPv4WriteIP(t, true, addr4)
|
||||||
|
})
|
||||||
})
|
})
|
||||||
})
|
})
|
||||||
// const addr6 = "2606:4700:4700::1001"
|
// const addr6 = "2606:4700:4700::1001"
|
||||||
|
|
@ -56,6 +62,9 @@ func TestPing(t *testing.T) {
|
||||||
t.Run("read-ip", func(t *testing.T) {
|
t.Run("read-ip", func(t *testing.T) {
|
||||||
testPingIPv6ReadIP(t, false, addr6)
|
testPingIPv6ReadIP(t, false, addr6)
|
||||||
})
|
})
|
||||||
|
t.Run("write-ip", func(t *testing.T) {
|
||||||
|
testPingIPv6WriteIP(t, false, addr6)
|
||||||
|
})
|
||||||
})
|
})
|
||||||
t.Run("privileged", func(t *testing.T) {
|
t.Run("privileged", func(t *testing.T) {
|
||||||
if runtime.GOOS != "windows" && os.Getuid() != 0 {
|
if runtime.GOOS != "windows" && os.Getuid() != 0 {
|
||||||
|
|
@ -67,6 +76,9 @@ func TestPing(t *testing.T) {
|
||||||
t.Run("read-ip", func(t *testing.T) {
|
t.Run("read-ip", func(t *testing.T) {
|
||||||
testPingIPv6ReadIP(t, true, addr6)
|
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, header.ICMPv6EchoReply, icmpHdr.Type())
|
||||||
require.Equal(t, request.Ident(), icmpHdr.Ident())
|
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())
|
||||||
|
}
|
||||||
|
|
|
||||||
|
|
@ -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(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)
|
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)
|
icmpForwarder.SetLocalAddresses(t.inet4Address, t.inet6Address)
|
||||||
ipStack.SetTransportProtocolHandler(icmp.ProtocolNumber4, icmpForwarder.HandlePacket)
|
ipStack.SetTransportProtocolHandler(icmp.ProtocolNumber4, icmpForwarder.HandlePacket)
|
||||||
ipStack.SetTransportProtocolHandler(icmp.ProtocolNumber6, icmpForwarder.HandlePacket)
|
ipStack.SetTransportProtocolHandler(icmp.ProtocolNumber6, icmpForwarder.HandlePacket)
|
||||||
|
|
|
||||||
|
|
@ -18,6 +18,8 @@ import (
|
||||||
"github.com/sagernet/gvisor/pkg/tcpip/network/ipv6"
|
"github.com/sagernet/gvisor/pkg/tcpip/network/ipv6"
|
||||||
"github.com/sagernet/gvisor/pkg/tcpip/stack"
|
"github.com/sagernet/gvisor/pkg/tcpip/stack"
|
||||||
"github.com/sagernet/sing/common/buf"
|
"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"
|
M "github.com/sagernet/sing/common/metadata"
|
||||||
N "github.com/sagernet/sing/common/network"
|
N "github.com/sagernet/sing/common/network"
|
||||||
)
|
)
|
||||||
|
|
@ -25,6 +27,7 @@ import (
|
||||||
type ICMPForwarder struct {
|
type ICMPForwarder struct {
|
||||||
ctx context.Context
|
ctx context.Context
|
||||||
stack *stack.Stack
|
stack *stack.Stack
|
||||||
|
logger logger.Logger
|
||||||
inet4Address netip.Addr
|
inet4Address netip.Addr
|
||||||
inet6Address netip.Addr
|
inet6Address netip.Addr
|
||||||
handler Handler
|
handler Handler
|
||||||
|
|
@ -34,12 +37,14 @@ type ICMPForwarder struct {
|
||||||
func NewICMPForwarder(
|
func NewICMPForwarder(
|
||||||
ctx context.Context,
|
ctx context.Context,
|
||||||
stack *stack.Stack,
|
stack *stack.Stack,
|
||||||
|
logger logger.Logger,
|
||||||
handler Handler,
|
handler Handler,
|
||||||
timeout time.Duration,
|
timeout time.Duration,
|
||||||
) *ICMPForwarder {
|
) *ICMPForwarder {
|
||||||
return &ICMPForwarder{
|
return &ICMPForwarder{
|
||||||
ctx: ctx,
|
ctx: ctx,
|
||||||
stack: stack,
|
stack: stack,
|
||||||
|
logger: logger,
|
||||||
handler: handler,
|
handler: handler,
|
||||||
mapping: NewDirectRouteMapping(timeout),
|
mapping: NewDirectRouteMapping(timeout),
|
||||||
}
|
}
|
||||||
|
|
@ -81,8 +86,10 @@ func (f *ICMPForwarder) HandlePacket(id stack.TransportEndpointID, pkt *stack.Pa
|
||||||
return true
|
return true
|
||||||
}
|
}
|
||||||
if action != nil {
|
if action != nil {
|
||||||
// TODO: handle error
|
err = icmpWritePacketBuffer(action, pkt)
|
||||||
_ = icmpWritePacketBuffer(action, pkt)
|
if err != nil {
|
||||||
|
f.logger.Error(E.Cause(err, "write ICMPv4 echo request"))
|
||||||
|
}
|
||||||
return true
|
return true
|
||||||
}
|
}
|
||||||
}
|
}
|
||||||
|
|
@ -95,7 +102,7 @@ func (f *ICMPForwarder) HandlePacket(id stack.TransportEndpointID, pkt *stack.Pa
|
||||||
ipHdr.SetChecksum(^ipHdr.CalculateChecksum())
|
ipHdr.SetChecksum(^ipHdr.CalculateChecksum())
|
||||||
outgoingEP, gErr := f.stack.GetNetworkEndpoint(DefaultNIC, header.IPv4ProtocolNumber)
|
outgoingEP, gErr := f.stack.GetNetworkEndpoint(DefaultNIC, header.IPv4ProtocolNumber)
|
||||||
if gErr != nil {
|
if gErr != nil {
|
||||||
// TODO: log error
|
f.logger.Error(E.Cause(gonet.TranslateNetstackError(gErr), "get IPv4 network endpoint"))
|
||||||
return true
|
return true
|
||||||
}
|
}
|
||||||
route, gErr := f.stack.FindRoute(
|
route, gErr := f.stack.FindRoute(
|
||||||
|
|
@ -106,7 +113,7 @@ func (f *ICMPForwarder) HandlePacket(id stack.TransportEndpointID, pkt *stack.Pa
|
||||||
false,
|
false,
|
||||||
)
|
)
|
||||||
if gErr != nil {
|
if gErr != nil {
|
||||||
// TODO: log error
|
f.logger.Error(E.Cause(gonet.TranslateNetstackError(gErr), "find IPv4 route"))
|
||||||
return true
|
return true
|
||||||
}
|
}
|
||||||
defer route.Release()
|
defer route.Release()
|
||||||
|
|
@ -142,9 +149,11 @@ func (f *ICMPForwarder) HandlePacket(id stack.TransportEndpointID, pkt *stack.Pa
|
||||||
return true
|
return true
|
||||||
}
|
}
|
||||||
if action != nil {
|
if action != nil {
|
||||||
// TODO: handle error
|
|
||||||
pkt.IncRef()
|
pkt.IncRef()
|
||||||
_ = icmpWritePacketBuffer(action, pkt)
|
err = icmpWritePacketBuffer(action, pkt)
|
||||||
|
if err != nil {
|
||||||
|
f.logger.Error(E.Cause(err, "write ICMPv6 echo request"))
|
||||||
|
}
|
||||||
return true
|
return true
|
||||||
}
|
}
|
||||||
}
|
}
|
||||||
|
|
@ -161,7 +170,7 @@ func (f *ICMPForwarder) HandlePacket(id stack.TransportEndpointID, pkt *stack.Pa
|
||||||
}))
|
}))
|
||||||
outgoingEP, gErr := f.stack.GetNetworkEndpoint(DefaultNIC, header.IPv4ProtocolNumber)
|
outgoingEP, gErr := f.stack.GetNetworkEndpoint(DefaultNIC, header.IPv4ProtocolNumber)
|
||||||
if gErr != nil {
|
if gErr != nil {
|
||||||
// TODO: log error
|
f.logger.Error(E.Cause(gonet.TranslateNetstackError(gErr), "get IPv6 network endpoint"))
|
||||||
return true
|
return true
|
||||||
}
|
}
|
||||||
route, gErr := f.stack.FindRoute(
|
route, gErr := f.stack.FindRoute(
|
||||||
|
|
@ -172,7 +181,7 @@ func (f *ICMPForwarder) HandlePacket(id stack.TransportEndpointID, pkt *stack.Pa
|
||||||
false,
|
false,
|
||||||
)
|
)
|
||||||
if gErr != nil {
|
if gErr != nil {
|
||||||
// TODO: log error
|
f.logger.Error(E.Cause(gonet.TranslateNetstackError(gErr), "find IPv6 route"))
|
||||||
return true
|
return true
|
||||||
}
|
}
|
||||||
defer route.Release()
|
defer route.Release()
|
||||||
|
|
|
||||||
Loading…
Add table
Add a link
Reference in a new issue