diff --git a/ping/destination.go b/ping/destination.go index 8648ecc..e16763b 100644 --- a/ping/destination.go +++ b/ping/destination.go @@ -105,23 +105,29 @@ func (d *Destination) loopRead() { } icmpHdr := header.ICMPv4(ipHdr.Payload()) if d.needFilter() { - if icmpHdr.Type() != header.ICMPv4EchoReply { - continue - } - var requestExists bool - request := pingRequest{Source: ipHdr.DestinationAddr(), Destination: ipHdr.SourceAddr(), Identifier: icmpHdr.Ident(), Sequence: icmpHdr.Sequence()} - d.requestAccess.Lock() - _, loaded := d.requests[request] - if loaded { - requestExists = true - delete(d.requests, request) - } - d.requestAccess.Unlock() - if !requestExists { + switch icmpHdr.Type() { + case header.ICMPv4EchoReply: + request := pingRequest{Source: ipHdr.DestinationAddr(), Destination: ipHdr.SourceAddr(), Identifier: icmpHdr.Ident(), Sequence: icmpHdr.Sequence()} + d.requestAccess.Lock() + _, loaded := d.requests[request] + if loaded { + delete(d.requests, request) + } + d.requestAccess.Unlock() + if !loaded { + continue + } + d.logger.TraceContext(d.ctx, "read ICMPv4 echo reply from ", ipHdr.SourceAddr(), " to ", ipHdr.DestinationAddr(), " id ", icmpHdr.Ident(), " seq ", icmpHdr.Sequence()) + case header.ICMPv4TimeExceeded, header.ICMPv4DstUnreachable: + if !d.rewriteICMPv4Error(ipHdr, icmpHdr) { + continue + } + default: continue } + } else { + d.logger.TraceContext(d.ctx, "read ICMPv4 echo reply from ", ipHdr.SourceAddr(), " to ", ipHdr.DestinationAddr(), " id ", icmpHdr.Ident(), " seq ", icmpHdr.Sequence()) } - d.logger.TraceContext(d.ctx, "read ICMPv4 echo reply from ", ipHdr.SourceAddr(), " to ", ipHdr.DestinationAddr(), " id ", icmpHdr.Ident(), " seq ", icmpHdr.Sequence()) } else { ipHdr := header.IPv6(buffer.Bytes()) if !ipHdr.IsValid(buffer.Len()) { @@ -191,6 +197,45 @@ func (d *Destination) WritePacket(packet *buf.Buffer) error { return d.conn.WriteIP(packet) } +func (d *Destination) rewriteICMPv4Error(ipHdr header.IPv4, icmpHdr header.ICMPv4) bool { + inner := icmpHdr.Payload() + if len(inner) < header.IPv4MinimumSize { + return false + } + innerIPHdr := header.IPv4(inner) + headerLen := int(innerIPHdr.HeaderLength()) + if headerLen < header.IPv4MinimumSize || len(inner) < headerLen+header.ICMPv4MinimumSize { + return false + } + if innerIPHdr.TransportProtocol() != header.ICMPv4ProtocolNumber { + return false + } + innerICMP := header.ICMPv4(inner[headerLen:]) + if innerICMP.Type() != header.ICMPv4Echo { + return false + } + originalIdent := ^innerICMP.Ident() + request := pingRequest{ + Source: ipHdr.DestinationAddr(), + Destination: innerIPHdr.DestinationAddr(), + Identifier: originalIdent, + Sequence: innerICMP.Sequence(), + } + d.requestAccess.Lock() + _, loaded := d.requests[request] + d.requestAccess.Unlock() + if !loaded { + return false + } + innerICMP.SetIdent(originalIdent) + innerICMP.SetChecksum(header.ICMPv4Checksum(innerICMP, 0)) + innerIPHdr.SetSourceAddr(ipHdr.DestinationAddr()) + innerIPHdr.SetChecksum(^innerIPHdr.CalculateChecksum()) + icmpHdr.SetChecksum(header.ICMPv4Checksum(icmpHdr, 0)) + d.logger.TraceContext(d.ctx, "read ICMPv4 error type ", int(icmpHdr.Type()), " from ", ipHdr.SourceAddr(), " seq ", innerICMP.Sequence()) + return true +} + func (d *Destination) needFilter() bool { return !d.conn.isLinuxUnprivileged() } diff --git a/ping/ping.go b/ping/ping.go index 248987c..d855424 100644 --- a/ping/ping.go +++ b/ping/ping.go @@ -158,9 +158,19 @@ func (c *Conn) ReadIP(buffer *buf.Buffer) error { }) } } else { - _, err := buffer.ReadOnceFrom(c.conn) - if err != nil { - return err + if runtime.GOOS == "linux" || runtime.GOOS == "android" || runtime.GOOS == "windows" { + // An unconnected SOCK_RAW IPv4 socket delivers the full packet including the IP + // header via ReadMsgIP, whereas ReadFrom strips it. + n, _, _, err := c.readMsg(buffer.FreeBytes(), nil) + if err != nil { + return err + } + buffer.Truncate(n) + } else { + _, err := buffer.ReadOnceFrom(c.conn) + if err != nil { + return err + } } if !c.destination.Is6() { ipHdr := header.IPv4(buffer.Bytes()) @@ -177,10 +187,12 @@ func (c *Conn) ReadIP(buffer *buf.Buffer) error { ipHdr.SetDestinationAddr(c.source.Load()) ipHdr.SetChecksum(^ipHdr.CalculateChecksum()) icmpHdr := header.ICMPv4(ipHdr.Payload()) - if !c.isLinuxUnprivileged() { - icmpHdr.SetIdent(^icmpHdr.Ident()) + if icmpHdr.Type() == header.ICMPv4EchoReply { + if !c.isLinuxUnprivileged() { + icmpHdr.SetIdent(^icmpHdr.Ident()) + } + icmpHdr.SetChecksum(header.ICMPv4Checksum(icmpHdr, 0)) } - icmpHdr.SetChecksum(header.ICMPv4Checksum(icmpHdr, 0)) } else { ipHdr := header.IPv6(buffer.Bytes()) if !ipHdr.IsValid(buffer.Len()) { @@ -202,27 +214,41 @@ func (c *Conn) ReadIP(buffer *buf.Buffer) error { } func (c *Conn) ReadICMP(buffer *buf.Buffer) error { + if !c.isLinuxUnprivileged() && !c.destination.Is6() { + if runtime.GOOS == "linux" || runtime.GOOS == "android" || runtime.GOOS == "windows" { + // An unconnected SOCK_RAW IPv4 socket delivers the full packet including the IP + // header via ReadMsgIP, whereas ReadFrom strips it. + n, _, _, err := c.readMsg(buffer.FreeBytes(), nil) + if err != nil { + return err + } + buffer.Truncate(n) + } else { + _, err := buffer.ReadOnceFrom(c.conn) + if err != nil { + return err + } + } + ipHdr := header.IPv4(buffer.Bytes()) + buffer.Advance(int(ipHdr.HeaderLength())) + + icmpHdr := header.ICMPv4(buffer.Bytes()) + icmpHdr.SetIdent(^icmpHdr.Ident()) + icmpHdr.SetChecksum(header.ICMPv4Checksum(icmpHdr, 0)) + return nil + } _, err := buffer.ReadOnceFrom(c.conn) if err != nil { return err } - if !c.isLinuxUnprivileged() { - if !c.destination.Is6() { - ipHdr := header.IPv4(buffer.Bytes()) - buffer.Advance(int(ipHdr.HeaderLength())) - - icmpHdr := header.ICMPv4(buffer.Bytes()) - icmpHdr.SetIdent(^icmpHdr.Ident()) - icmpHdr.SetChecksum(header.ICMPv4Checksum(icmpHdr, 0)) - } else { - icmpHdr := header.ICMPv6(buffer.Bytes()) - icmpHdr.SetIdent(^icmpHdr.Ident()) - icmpHdr.SetChecksum(header.ICMPv6Checksum(header.ICMPv6ChecksumParams{ - Header: icmpHdr, - Src: c.destination.AsSlice(), - Dst: c.source.Load().AsSlice(), - })) - } + if c.destination.Is6() && !c.isLinuxUnprivileged() { + icmpHdr := header.ICMPv6(buffer.Bytes()) + icmpHdr.SetIdent(^icmpHdr.Ident()) + icmpHdr.SetChecksum(header.ICMPv6Checksum(header.ICMPv6ChecksumParams{ + Header: icmpHdr, + Src: c.destination.AsSlice(), + Dst: c.source.Load().AsSlice(), + })) } return nil } @@ -232,6 +258,10 @@ 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())) + if err != nil { + return err + } icmpHdr := header.ICMPv4(ipHdr.Payload()) icmpHdr.SetIdent(^icmpHdr.Ident()) icmpHdr.SetChecksum(header.ICMPv4Checksum(icmpHdr, 0)) @@ -241,6 +271,10 @@ 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())) + if err != nil { + return err + } icmpHdr := header.ICMPv6(ipHdr.Payload()) icmpHdr.SetIdent(^icmpHdr.Ident()) icmpHdr.SetChecksum(header.ICMPv6Checksum(header.ICMPv6ChecksumParams{ diff --git a/ping/socket_unix.go b/ping/socket_unix.go index 1eec10a..9128239 100644 --- a/ping/socket_unix.go +++ b/ping/socket_unix.go @@ -87,24 +87,33 @@ func connect(privileged bool, controlFunc control.Func, destination netip.Addr) return nil, err } - if runtime.GOOS == "darwin" && !privileged { - // When running in NetworkExtension on macOS, write to connected socket results in EPIPE. + useUnconnected := (runtime.GOOS == "darwin" && !privileged) || + ((runtime.GOOS == "linux" || runtime.GOOS == "android") && privileged) + if useUnconnected { + // A connected ICMP socket only receives messages whose source is the connected peer, + // so the Time Exceeded replies that transit routers send for traceroute never reach it. + // Additionally, on macOS NetworkExtension, writing to a connected socket returns EPIPE. var packetConn net.PacketConn packetConn, err = net.FilePacketConn(file) if err != nil { return nil, err } - return bufio.NewBindPacketConn(packetConn, M.SocksaddrFrom(destination, 0).UDPAddr()), nil - } else { - err = unix.Connect(fd, M.AddrPortToSockaddr(netip.AddrPortFrom(destination, 0))) - if err != nil { - return nil, err + var writeTarget net.Addr + if privileged { + writeTarget = M.SocksaddrFrom(destination, 0).IPAddr() + } else { + writeTarget = M.SocksaddrFrom(destination, 0).UDPAddr() } - var conn net.Conn - conn, err = net.FileConn(file) - if err != nil { - return nil, err - } - return conn, nil + return bufio.NewBindPacketConn(packetConn, writeTarget), nil } + err = unix.Connect(fd, M.AddrPortToSockaddr(netip.AddrPortFrom(destination, 0))) + if err != nil { + return nil, err + } + var conn net.Conn + conn, err = net.FileConn(file) + if err != nil { + return nil, err + } + return conn, nil } diff --git a/ping/socket_windows.go b/ping/socket_windows.go index daafd18..332a251 100644 --- a/ping/socket_windows.go +++ b/ping/socket_windows.go @@ -1,30 +1,29 @@ package ping import ( + "context" "net" "net/netip" "syscall" + "github.com/sagernet/sing/common/bufio" "github.com/sagernet/sing/common/control" + M "github.com/sagernet/sing/common/metadata" "golang.org/x/sys/windows" ) func connect(privileged bool, controlFunc control.Func, destination netip.Addr) (net.Conn, error) { - var dialer net.Dialer - dialer.Control = controlFunc + var listenConfig net.ListenConfig + listenConfig.Control = controlFunc if destination.Is6() { - dialer.Control = control.Append(dialer.Control, func(network, address string, conn syscall.RawConn) error { + listenConfig.Control = control.Append(listenConfig.Control, func(network, address string, conn syscall.RawConn) error { return control.Raw(conn, func(fd uintptr) error { err := windows.SetsockoptInt(windows.Handle(fd), windows.IPPROTO_IPV6, IPV6_HOPLIMIT, 1) if err != nil { return err } - err = windows.SetsockoptInt(windows.Handle(fd), windows.IPPROTO_IPV6, IPV6_RECVTCLASS, 1) - if err != nil { - return err - } - return nil + return windows.SetsockoptInt(windows.Handle(fd), windows.IPPROTO_IPV6, IPV6_RECVTCLASS, 1) }) }) } @@ -34,5 +33,11 @@ func connect(privileged bool, controlFunc control.Func, destination netip.Addr) } else { network = "ip6:ipv6-icmp" } - return dialer.Dial(network, destination.String()) + // A connected raw socket only receives messages from the connected peer, so transit routers' + // Time Exceeded replies needed by traceroute never arrive. + packetConn, err := listenConfig.ListenPacket(context.Background(), network, "") + if err != nil { + return nil, err + } + return bufio.NewBindPacketConn(packetConn, M.SocksaddrFrom(destination, 0).IPAddr()), nil } diff --git a/tun_darwin.go b/tun_darwin.go index f5f7aff..3131295 100644 --- a/tun_darwin.go +++ b/tun_darwin.go @@ -10,7 +10,7 @@ import ( "unsafe" "github.com/sagernet/sing-tun/internal/gtcpip/header" - "github.com/sagernet/sing-tun/internal/rawfile_darwin" + rawfile "github.com/sagernet/sing-tun/internal/rawfile_darwin" "github.com/sagernet/sing-tun/internal/stopfd_darwin" "github.com/sagernet/sing/common" "github.com/sagernet/sing/common/buf"