mirror of
https://github.com/XTLS/Xray-core.git
synced 2026-09-29 02:18:04 +03:00
Completes https://github.com/XTLS/Xray-core/pull/6807 and https://github.com/XTLS/Xray-core/pull/6810
329 lines
8.9 KiB
Go
329 lines
8.9 KiB
Go
package masque
|
|
|
|
import (
|
|
"bytes"
|
|
"context"
|
|
"io"
|
|
"net/netip"
|
|
"os"
|
|
"sync"
|
|
"testing"
|
|
"time"
|
|
|
|
"github.com/stretchr/testify/require"
|
|
"github.com/xtls/xray-core/common/buf"
|
|
"github.com/xtls/xray-core/common/net"
|
|
"github.com/xtls/xray-core/common/protocol"
|
|
"golang.zx2c4.com/wireguard/tun"
|
|
)
|
|
|
|
type fakeTunnelConn struct {
|
|
mu sync.Mutex
|
|
reads chan []byte
|
|
written [][]byte
|
|
closed bool
|
|
stall chan struct{}
|
|
}
|
|
|
|
func newFakeTunnelConn() *fakeTunnelConn {
|
|
return &fakeTunnelConn{reads: make(chan []byte, 16)}
|
|
}
|
|
|
|
func (c *fakeTunnelConn) Read(b []byte) (int, error) {
|
|
p, ok := <-c.reads
|
|
if !ok {
|
|
return 0, io.EOF
|
|
}
|
|
return copy(b, p), nil
|
|
}
|
|
|
|
func (c *fakeTunnelConn) Write(b []byte) (int, error) {
|
|
if c.stall != nil {
|
|
<-c.stall
|
|
}
|
|
c.mu.Lock()
|
|
defer c.mu.Unlock()
|
|
c.written = append(c.written, bytes.Clone(b))
|
|
return len(b), nil
|
|
}
|
|
|
|
func (c *fakeTunnelConn) Close() error {
|
|
c.mu.Lock()
|
|
defer c.mu.Unlock()
|
|
if !c.closed {
|
|
c.closed = true
|
|
close(c.reads)
|
|
}
|
|
return nil
|
|
}
|
|
|
|
func (c *fakeTunnelConn) packets() [][]byte {
|
|
c.mu.Lock()
|
|
defer c.mu.Unlock()
|
|
return c.written
|
|
}
|
|
|
|
func (c *fakeTunnelConn) isClosed() bool {
|
|
c.mu.Lock()
|
|
defer c.mu.Unlock()
|
|
return c.closed
|
|
}
|
|
|
|
func (c *fakeTunnelConn) LocalAddr() net.Addr { return &net.TCPAddr{} }
|
|
func (c *fakeTunnelConn) RemoteAddr() net.Addr { return &net.TCPAddr{} }
|
|
func (c *fakeTunnelConn) SetDeadline(t time.Time) error { return nil }
|
|
func (c *fakeTunnelConn) SetReadDeadline(t time.Time) error { return nil }
|
|
func (c *fakeTunnelConn) SetWriteDeadline(t time.Time) error { return nil }
|
|
|
|
type fakeDevice struct {
|
|
mu sync.Mutex
|
|
reads chan []byte
|
|
written [][]byte
|
|
closed bool
|
|
}
|
|
|
|
func (d *fakeDevice) File() *os.File { return nil }
|
|
func (d *fakeDevice) MTU() (int, error) { return 1280, nil }
|
|
func (d *fakeDevice) Name() (string, error) { return "fake", nil }
|
|
func (d *fakeDevice) Events() <-chan tun.Event { return nil }
|
|
func (d *fakeDevice) BatchSize() int { return 1 }
|
|
|
|
func (d *fakeDevice) Read(bufs [][]byte, sizes []int, offset int) (int, error) {
|
|
p, ok := <-d.reads
|
|
if !ok {
|
|
return 0, os.ErrClosed
|
|
}
|
|
sizes[0] = copy(bufs[0][offset:], p)
|
|
return 1, nil
|
|
}
|
|
|
|
func (d *fakeDevice) Write(bufs [][]byte, offset int) (int, error) {
|
|
d.mu.Lock()
|
|
defer d.mu.Unlock()
|
|
for _, b := range bufs {
|
|
d.written = append(d.written, bytes.Clone(b[offset:]))
|
|
}
|
|
return len(bufs), nil
|
|
}
|
|
|
|
func (d *fakeDevice) Close() error {
|
|
d.mu.Lock()
|
|
defer d.mu.Unlock()
|
|
if !d.closed {
|
|
d.closed = true
|
|
close(d.reads)
|
|
}
|
|
return nil
|
|
}
|
|
|
|
func (d *fakeDevice) packets() [][]byte {
|
|
d.mu.Lock()
|
|
defer d.mu.Unlock()
|
|
return d.written
|
|
}
|
|
|
|
func ipPacket(src, dst string) []byte {
|
|
s, d := netip.MustParseAddr(src), netip.MustParseAddr(dst)
|
|
if s.Is4() {
|
|
b := make([]byte, 20)
|
|
b[0] = 0x45
|
|
b[8] = 64
|
|
copy(b[12:16], s.AsSlice())
|
|
copy(b[16:20], d.AsSlice())
|
|
return b
|
|
}
|
|
b := make([]byte, 40)
|
|
b[0] = 0x60
|
|
b[7] = 64
|
|
copy(b[8:24], s.AsSlice())
|
|
copy(b[24:40], d.AsSlice())
|
|
return b
|
|
}
|
|
|
|
func newTestServer(t *testing.T) (*Server, *fakeDevice) {
|
|
t.Helper()
|
|
pool4, err := newAddressPool(netip.MustParsePrefix("10.14.0.1/24"))
|
|
require.NoError(t, err)
|
|
pool6, err := newAddressPool(netip.MustParsePrefix("fd14::1/64"))
|
|
require.NoError(t, err)
|
|
dev := &fakeDevice{reads: make(chan []byte, tunnelQueueSize*2)}
|
|
s := &Server{
|
|
mtu: 1280,
|
|
dev: dev,
|
|
pools: []*addressPool{pool4, pool6},
|
|
local: []netip.Addr{netip.MustParseAddr("10.14.0.1"), netip.MustParseAddr("fd14::1")},
|
|
tunnels: make(map[netip.Addr]*serverTunnel),
|
|
}
|
|
return s, dev
|
|
}
|
|
|
|
func addTunnel(t *testing.T, s *Server) (*serverTunnel, *fakeTunnelConn) {
|
|
t.Helper()
|
|
return addUserTunnel(t, s, &protocol.MemoryUser{})
|
|
}
|
|
|
|
func addUserTunnel(t *testing.T, s *Server, user *protocol.MemoryUser) (*serverTunnel, *fakeTunnelConn) {
|
|
t.Helper()
|
|
conn := newFakeTunnelConn()
|
|
tunnel := newServerTunnel(conn, user)
|
|
for _, pool := range s.pools {
|
|
addr, ok := pool.allocate()
|
|
require.True(t, ok)
|
|
tunnel.addrs = append(tunnel.addrs, addr)
|
|
}
|
|
require.True(t, s.register(tunnel))
|
|
go s.writeToTunnel(tunnel)
|
|
t.Cleanup(tunnel.close)
|
|
return tunnel, conn
|
|
}
|
|
|
|
func TestServerRoutesTunnelPackets(t *testing.T) {
|
|
s, dev := newTestServer(t)
|
|
a, aConn := addTunnel(t, s)
|
|
b, bConn := addTunnel(t, s)
|
|
require.Equal(t, []netip.Addr{netip.MustParseAddr("10.14.0.2"), netip.MustParseAddr("fd14::2")}, a.addrs)
|
|
require.Equal(t, []netip.Addr{netip.MustParseAddr("10.14.0.3"), netip.MustParseAddr("fd14::3")}, b.addrs)
|
|
|
|
toB := ipPacket("10.14.0.2", "10.14.0.3")
|
|
toB6 := ipPacket("fd14::2", "fd14::3")
|
|
toServer := ipPacket("10.14.0.2", "10.14.0.1")
|
|
toInternet := ipPacket("fd14::2", "2001:db8::1")
|
|
for _, p := range [][]byte{
|
|
toB,
|
|
toB6,
|
|
ipPacket("10.14.0.2", "10.14.0.9"),
|
|
ipPacket("fd14::2", "fd14::99"),
|
|
ipPacket("fd14::2", "fe80::1"),
|
|
ipPacket("fd14::2", "ff02::1"),
|
|
ipPacket("10.14.0.2", "224.0.0.251"),
|
|
ipPacket("10.14.0.2", "10.14.0.2"),
|
|
toServer,
|
|
toInternet,
|
|
} {
|
|
aConn.reads <- p
|
|
}
|
|
aConn.Close()
|
|
require.NoError(t, s.readFromTunnel(a))
|
|
|
|
require.Eventually(t, func() bool { return len(bConn.packets()) == 2 }, time.Second, time.Millisecond)
|
|
require.Equal(t, [][]byte{toB, toB6}, bConn.packets())
|
|
require.Equal(t, [][]byte{toServer, toInternet}, dev.packets())
|
|
require.Empty(t, aConn.packets())
|
|
}
|
|
|
|
func TestServerRoutesStackPackets(t *testing.T) {
|
|
s, dev := newTestServer(t)
|
|
_, aConn := addTunnel(t, s)
|
|
_, bConn := addTunnel(t, s)
|
|
require.NoError(t, s.Start())
|
|
|
|
toA := ipPacket("192.0.2.1", "10.14.0.2")
|
|
toB := ipPacket("2001:db8::1", "fd14::3")
|
|
dev.reads <- toA
|
|
dev.reads <- ipPacket("192.0.2.1", "10.14.0.9")
|
|
dev.reads <- toB
|
|
require.Eventually(t, func() bool {
|
|
return len(aConn.packets()) == 1 && len(bConn.packets()) == 1
|
|
}, time.Second, time.Millisecond)
|
|
require.Equal(t, [][]byte{toA}, aConn.packets())
|
|
require.Equal(t, [][]byte{toB}, bConn.packets())
|
|
|
|
require.NoError(t, s.Close())
|
|
require.True(t, aConn.isClosed())
|
|
require.True(t, bConn.isClosed())
|
|
require.False(t, s.register(&serverTunnel{}))
|
|
}
|
|
|
|
func TestServerSlowTunnelDoesNotBlockOthers(t *testing.T) {
|
|
s, dev := newTestServer(t)
|
|
_, aConn := addTunnel(t, s)
|
|
_, bConn := addTunnel(t, s)
|
|
aConn.stall = make(chan struct{})
|
|
defer close(aConn.stall)
|
|
require.NoError(t, s.Start())
|
|
defer s.Close()
|
|
|
|
for range tunnelQueueSize + 10 {
|
|
dev.reads <- ipPacket("192.0.2.1", "10.14.0.2")
|
|
}
|
|
toB := ipPacket("192.0.2.1", "10.14.0.3")
|
|
dev.reads <- toB
|
|
require.Eventually(t, func() bool { return len(bConn.packets()) == 1 }, time.Second, time.Millisecond)
|
|
require.Equal(t, [][]byte{toB}, bConn.packets())
|
|
}
|
|
|
|
func TestServerClosesTunnelConnections(t *testing.T) {
|
|
s, _ := newTestServer(t)
|
|
a, _ := addTunnel(t, s)
|
|
conn := newFakeTunnelConn()
|
|
require.True(t, a.track(conn))
|
|
other := newFakeTunnelConn()
|
|
require.True(t, a.track(other))
|
|
a.untrack(other)
|
|
|
|
s.release(a)
|
|
require.True(t, conn.isClosed())
|
|
require.False(t, other.isClosed())
|
|
require.False(t, a.track(newFakeTunnelConn()))
|
|
require.False(t, a.send(buf.New()))
|
|
}
|
|
|
|
func TestServerReleasesAddresses(t *testing.T) {
|
|
s, _ := newTestServer(t)
|
|
a, _ := addTunnel(t, s)
|
|
s.release(a)
|
|
require.Nil(t, s.lookup(netip.MustParseAddr("10.14.0.2")))
|
|
b, _ := addTunnel(t, s)
|
|
require.Equal(t, []netip.Addr{netip.MustParseAddr("10.14.0.3"), netip.MustParseAddr("fd14::3")}, b.addrs)
|
|
for range 250 {
|
|
addTunnel(t, s)
|
|
}
|
|
c, _ := addTunnel(t, s)
|
|
require.Equal(t, netip.MustParseAddr("10.14.0.254"), c.addrs[0])
|
|
addr, ok := s.pools[0].allocate()
|
|
require.True(t, ok)
|
|
require.Equal(t, netip.MustParseAddr("10.14.0.2"), addr)
|
|
_, ok = s.pools[0].allocate()
|
|
require.False(t, ok)
|
|
}
|
|
|
|
func TestServerRemoveUserClosesTunnels(t *testing.T) {
|
|
s, _ := newTestServer(t)
|
|
s.validator = newValidator()
|
|
alice := &protocol.MemoryUser{Email: "a@example.com", Account: &MemoryAccount{Password: "p"}}
|
|
bob := &protocol.MemoryUser{Email: "b@example.com", Account: &MemoryAccount{Password: "p"}}
|
|
require.NoError(t, s.AddUser(context.Background(), alice))
|
|
require.NoError(t, s.AddUser(context.Background(), bob))
|
|
_, aConn := addUserTunnel(t, s, alice)
|
|
_, bConn := addUserTunnel(t, s, bob)
|
|
|
|
require.NoError(t, s.RemoveUser(context.Background(), "a@example.com"))
|
|
require.True(t, aConn.isClosed())
|
|
require.False(t, bConn.isClosed())
|
|
require.Error(t, s.RemoveUser(context.Background(), "a@example.com"))
|
|
require.Nil(t, s.validator.get("a@example.com", "p"))
|
|
require.Equal(t, bob, s.validator.get("b@example.com", "p"))
|
|
}
|
|
|
|
func TestPacketDestination(t *testing.T) {
|
|
v4 := make([]byte, 20)
|
|
v4[0] = 0x45
|
|
copy(v4[16:20], []byte{192, 0, 2, 1})
|
|
addr, ok := packetDestination(v4)
|
|
require.True(t, ok)
|
|
require.Equal(t, netip.MustParseAddr("192.0.2.1"), addr)
|
|
|
|
v6 := make([]byte, 40)
|
|
v6[0] = 0x60
|
|
dst := netip.MustParseAddr("2001:db8::1").As16()
|
|
copy(v6[24:40], dst[:])
|
|
addr, ok = packetDestination(v6)
|
|
require.True(t, ok)
|
|
require.Equal(t, netip.MustParseAddr("2001:db8::1"), addr)
|
|
|
|
for _, b := range [][]byte{nil, v4[:19], v6[:39], {0x50}} {
|
|
_, ok = packetDestination(b)
|
|
require.False(t, ok)
|
|
}
|
|
}
|