Files
Xray-core/proxy/wireguard/icmp_test.go
T

177 lines
5.2 KiB
Go

package wireguard
import (
"bytes"
"net/netip"
"testing"
"time"
"github.com/xtls/xray-core/common/net"
"gvisor.dev/gvisor/pkg/tcpip"
"gvisor.dev/gvisor/pkg/tcpip/checksum"
"gvisor.dev/gvisor/pkg/tcpip/header"
)
func newICMPTestStack(t *testing.T) *netTun {
t.Helper()
dev, _, gstack, err := CreateNetTUN([]netip.Addr{
netip.MustParseAddr("10.66.0.1"),
netip.MustParseAddr("fd00::1"),
}, nil, 1420, false)
if err != nil {
t.Fatal(err)
}
t.Cleanup(func() { dev.Close() })
CreateForwarder(gstack, func(conn net.Conn, dest net.Destination) { conn.Close() })
CreateICMPEchoResponder(gstack)
return dev.(*netTun)
}
// startReader must run before the request is written: the stack may answer
// synchronously inside Write, and netTun hands packets over an unbuffered channel.
func startReader(dev *netTun) <-chan []byte {
got := make(chan []byte, 1)
go func() {
buf := make([]byte, 2048)
sizes := make([]int, 1)
if _, err := dev.Read([][]byte{buf}, sizes, 0); err == nil {
got <- buf[:sizes[0]]
}
}()
return got
}
func awaitPacket(t *testing.T, got <-chan []byte) []byte {
t.Helper()
select {
case p := <-got:
return p
case <-time.After(2 * time.Second):
t.Fatal("no echo reply from the stack")
return nil
}
}
func TestICMPv4EchoReply(t *testing.T) {
dev := newICMPTestStack(t)
src := tcpip.AddrFrom4([4]byte{10, 66, 0, 2})
dst := tcpip.AddrFrom4([4]byte{1, 1, 1, 1})
payload := []byte("xray wireguard ping")
icmpMsg := make([]byte, header.ICMPv4MinimumSize+len(payload))
req := header.ICMPv4(icmpMsg)
req.SetType(header.ICMPv4Echo)
req.SetIdent(0x1234)
req.SetSequence(7)
copy(req.Payload(), payload)
req.SetChecksum(header.ICMPv4Checksum(req[:header.ICMPv4MinimumSize], checksum.Checksum(payload, 0)))
pkt := make([]byte, header.IPv4MinimumSize+len(icmpMsg))
ip := header.IPv4(pkt)
ip.Encode(&header.IPv4Fields{
TotalLength: uint16(len(pkt)),
TTL: 64,
Protocol: uint8(header.ICMPv4ProtocolNumber),
SrcAddr: src,
DstAddr: dst,
})
ip.SetChecksum(^ip.CalculateChecksum())
copy(pkt[header.IPv4MinimumSize:], icmpMsg)
got := startReader(dev)
if _, err := dev.Write([][]byte{pkt}, 0); err != nil {
t.Fatal(err)
}
reply := header.IPv4(awaitPacket(t, got))
if !reply.IsValid(len(reply)) {
t.Fatal("invalid ipv4 reply")
}
if reply.SourceAddress() != dst || reply.DestinationAddress() != src {
t.Fatalf("reply addresses %v -> %v, want %v -> %v", reply.SourceAddress(), reply.DestinationAddress(), dst, src)
}
if reply.TransportProtocol() != header.ICMPv4ProtocolNumber {
t.Fatalf("reply protocol %v, want icmpv4", reply.TransportProtocol())
}
echo := header.ICMPv4(reply.Payload())
if echo.Type() != header.ICMPv4EchoReply {
t.Fatalf("reply type %v, want echo reply", echo.Type())
}
if echo.Ident() != 0x1234 || echo.Sequence() != 7 {
t.Fatalf("reply ident/seq %#x/%d, want 0x1234/7", echo.Ident(), echo.Sequence())
}
if !bytes.Equal(echo.Payload(), payload) {
t.Fatalf("reply payload %q, want %q", echo.Payload(), payload)
}
if checksum.Checksum(echo, 0) != 0xffff {
t.Fatal("bad icmpv4 checksum")
}
}
func TestICMPv6EchoReply(t *testing.T) {
dev := newICMPTestStack(t)
src := tcpip.AddrFrom16([16]byte{0xfd, 15: 2})
dst := tcpip.AddrFrom16([16]byte{0x26, 0x06, 0x47, 0x00, 0x47, 0x00, 15: 0x11})
payload := []byte("xray wireguard ping6")
icmpMsg := make([]byte, header.ICMPv6MinimumSize+len(payload))
req := header.ICMPv6(icmpMsg)
req.SetType(header.ICMPv6EchoRequest)
req.SetIdent(0x4321)
req.SetSequence(9)
copy(req.Payload(), payload)
req.SetChecksum(header.ICMPv6Checksum(header.ICMPv6ChecksumParams{
Header: req[:header.ICMPv6MinimumSize],
Src: src,
Dst: dst,
PayloadCsum: checksum.Checksum(payload, 0),
PayloadLen: len(payload),
}))
pkt := make([]byte, header.IPv6MinimumSize+len(icmpMsg))
ip := header.IPv6(pkt)
ip.Encode(&header.IPv6Fields{
PayloadLength: uint16(len(icmpMsg)),
TransportProtocol: header.ICMPv6ProtocolNumber,
HopLimit: 64,
SrcAddr: src,
DstAddr: dst,
})
copy(pkt[header.IPv6MinimumSize:], icmpMsg)
got := startReader(dev)
if _, err := dev.Write([][]byte{pkt}, 0); err != nil {
t.Fatal(err)
}
reply := header.IPv6(awaitPacket(t, got))
if !reply.IsValid(len(reply)) {
t.Fatal("invalid ipv6 reply")
}
if reply.SourceAddress() != dst || reply.DestinationAddress() != src {
t.Fatalf("reply addresses %v -> %v, want %v -> %v", reply.SourceAddress(), reply.DestinationAddress(), dst, src)
}
echo := header.ICMPv6(reply.Payload())
if echo.Type() != header.ICMPv6EchoReply {
t.Fatalf("reply type %v, want echo reply", echo.Type())
}
if echo.Ident() != 0x4321 || echo.Sequence() != 9 {
t.Fatalf("reply ident/seq %#x/%d, want 0x4321/9", echo.Ident(), echo.Sequence())
}
if !bytes.Equal(echo.Payload(), payload) {
t.Fatalf("reply payload %q, want %q", echo.Payload(), payload)
}
zeroed := header.ICMPv6(append([]byte(nil), echo[:header.ICMPv6MinimumSize]...))
zeroed.SetChecksum(0)
want := header.ICMPv6Checksum(header.ICMPv6ChecksumParams{
Header: zeroed,
Src: dst,
Dst: src,
PayloadCsum: checksum.Checksum(echo.Payload(), 0),
PayloadLen: len(echo.Payload()),
})
if echo.Checksum() != want {
t.Fatalf("icmpv6 checksum %#x, want %#x", echo.Checksum(), want)
}
}