mirror of
https://github.com/XTLS/Xray-core.git
synced 2026-10-05 15:23:33 +00:00
177 lines
5.2 KiB
Go
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)
|
|
}
|
|
}
|