mirror of
https://github.com/XTLS/Xray-core.git
synced 2026-09-29 20:46:59 +00:00
Proxy: Add MASQUE inbound (IETF CONNECT-IP server, RFC 9484) (#6844)
Completes https://github.com/XTLS/Xray-core/pull/6807 and https://github.com/XTLS/Xray-core/pull/6810
This commit is contained in:
@@ -0,0 +1,108 @@
|
||||
package masque
|
||||
|
||||
import (
|
||||
"crypto/subtle"
|
||||
"strings"
|
||||
"sync"
|
||||
|
||||
"github.com/xtls/xray-core/common/errors"
|
||||
"github.com/xtls/xray-core/common/protocol"
|
||||
"google.golang.org/protobuf/proto"
|
||||
)
|
||||
|
||||
func (a *Account) AsAccount() (protocol.Account, error) {
|
||||
return &MemoryAccount{Password: a.Password}, nil
|
||||
}
|
||||
|
||||
type MemoryAccount struct {
|
||||
Password string
|
||||
}
|
||||
|
||||
func (a *MemoryAccount) Equals(other protocol.Account) bool {
|
||||
b, ok := other.(*MemoryAccount)
|
||||
return ok && a.Password == b.Password
|
||||
}
|
||||
|
||||
func (a *MemoryAccount) ToProto() proto.Message {
|
||||
return &Account{Password: a.Password}
|
||||
}
|
||||
|
||||
type validator struct {
|
||||
mu sync.RWMutex
|
||||
users map[string]*protocol.MemoryUser
|
||||
}
|
||||
|
||||
func newValidator() *validator {
|
||||
return &validator{users: make(map[string]*protocol.MemoryUser)}
|
||||
}
|
||||
|
||||
func (v *validator) add(user *protocol.MemoryUser) error {
|
||||
account, ok := user.Account.(*MemoryAccount)
|
||||
if !ok {
|
||||
return errors.New("not a MASQUE account")
|
||||
}
|
||||
if user.Email == "" || strings.Contains(user.Email, ":") {
|
||||
return errors.New("invalid email ", user.Email)
|
||||
}
|
||||
if account.Password == "" {
|
||||
return errors.New("empty password for ", user.Email)
|
||||
}
|
||||
email := strings.ToLower(user.Email)
|
||||
v.mu.Lock()
|
||||
defer v.mu.Unlock()
|
||||
if _, found := v.users[email]; found {
|
||||
return errors.New("user ", user.Email, " already exists")
|
||||
}
|
||||
v.users[email] = user
|
||||
return nil
|
||||
}
|
||||
|
||||
func (v *validator) delByEmail(email string) (*protocol.MemoryUser, error) {
|
||||
key := strings.ToLower(email)
|
||||
v.mu.Lock()
|
||||
defer v.mu.Unlock()
|
||||
user, found := v.users[key]
|
||||
if !found {
|
||||
return nil, errors.New("user ", email, " not found")
|
||||
}
|
||||
delete(v.users, key)
|
||||
return user, nil
|
||||
}
|
||||
|
||||
func (v *validator) contains(user *protocol.MemoryUser) bool {
|
||||
v.mu.RLock()
|
||||
defer v.mu.RUnlock()
|
||||
return v.users[strings.ToLower(user.Email)] == user
|
||||
}
|
||||
|
||||
func (v *validator) get(email, password string) *protocol.MemoryUser {
|
||||
v.mu.RLock()
|
||||
user := v.users[strings.ToLower(email)]
|
||||
v.mu.RUnlock()
|
||||
if user == nil || subtle.ConstantTimeCompare([]byte(user.Account.(*MemoryAccount).Password), []byte(password)) != 1 {
|
||||
return nil
|
||||
}
|
||||
return user
|
||||
}
|
||||
|
||||
func (v *validator) getByEmail(email string) *protocol.MemoryUser {
|
||||
v.mu.RLock()
|
||||
defer v.mu.RUnlock()
|
||||
return v.users[strings.ToLower(email)]
|
||||
}
|
||||
|
||||
func (v *validator) getAll() []*protocol.MemoryUser {
|
||||
v.mu.RLock()
|
||||
defer v.mu.RUnlock()
|
||||
users := make([]*protocol.MemoryUser, 0, len(v.users))
|
||||
for _, user := range v.users {
|
||||
users = append(users, user)
|
||||
}
|
||||
return users
|
||||
}
|
||||
|
||||
func (v *validator) count() int64 {
|
||||
v.mu.RLock()
|
||||
defer v.mu.RUnlock()
|
||||
return int64(len(v.users))
|
||||
}
|
||||
@@ -0,0 +1,49 @@
|
||||
package masque
|
||||
|
||||
import (
|
||||
"testing"
|
||||
|
||||
"github.com/stretchr/testify/require"
|
||||
"github.com/xtls/xray-core/common/protocol"
|
||||
)
|
||||
|
||||
func TestValidator(t *testing.T) {
|
||||
v := newValidator()
|
||||
user := &protocol.MemoryUser{Email: "U@example.com", Account: &MemoryAccount{Password: "p"}}
|
||||
require.NoError(t, v.add(user))
|
||||
for _, u := range []*protocol.MemoryUser{
|
||||
{Email: "u@example.com", Account: &MemoryAccount{Password: "other"}},
|
||||
{Account: &MemoryAccount{Password: "p"}},
|
||||
{Email: "a:b", Account: &MemoryAccount{Password: "p"}},
|
||||
{Email: "b@example.com", Account: &MemoryAccount{}},
|
||||
} {
|
||||
require.Error(t, v.add(u), u.Email)
|
||||
}
|
||||
|
||||
require.Equal(t, user, v.get("u@example.com", "p"))
|
||||
require.Equal(t, user, v.get("U@EXAMPLE.COM", "p"))
|
||||
require.Nil(t, v.get("u@example.com", "x"))
|
||||
require.Nil(t, v.get("x@example.com", "p"))
|
||||
require.Nil(t, v.get("", ""))
|
||||
require.Equal(t, user, v.getByEmail("u@example.com"))
|
||||
require.Equal(t, []*protocol.MemoryUser{user}, v.getAll())
|
||||
require.Equal(t, int64(1), v.count())
|
||||
|
||||
require.True(t, v.contains(user))
|
||||
removed, err := v.delByEmail("u@EXAMPLE.com")
|
||||
require.NoError(t, err)
|
||||
require.Equal(t, user, removed)
|
||||
_, err = v.delByEmail("u@example.com")
|
||||
require.Error(t, err)
|
||||
require.False(t, v.contains(user))
|
||||
require.Nil(t, v.get("u@example.com", "p"))
|
||||
require.Zero(t, v.count())
|
||||
}
|
||||
|
||||
func TestAccount(t *testing.T) {
|
||||
account, err := (&Account{Password: "p"}).AsAccount()
|
||||
require.NoError(t, err)
|
||||
require.True(t, account.Equals(&MemoryAccount{Password: "p"}))
|
||||
require.False(t, account.Equals(&MemoryAccount{Password: "x"}))
|
||||
require.Equal(t, &Account{Password: "p"}, account.ToProto())
|
||||
}
|
||||
+125
-11
@@ -74,15 +74,125 @@ func (x *ClientConfig) GetRemoteDns() []string {
|
||||
return nil
|
||||
}
|
||||
|
||||
type Account struct {
|
||||
state protoimpl.MessageState `protogen:"open.v1"`
|
||||
Password string `protobuf:"bytes,1,opt,name=password,proto3" json:"password,omitempty"`
|
||||
unknownFields protoimpl.UnknownFields
|
||||
sizeCache protoimpl.SizeCache
|
||||
}
|
||||
|
||||
func (x *Account) Reset() {
|
||||
*x = Account{}
|
||||
mi := &file_proxy_masque_config_proto_msgTypes[1]
|
||||
ms := protoimpl.X.MessageStateOf(protoimpl.Pointer(x))
|
||||
ms.StoreMessageInfo(mi)
|
||||
}
|
||||
|
||||
func (x *Account) String() string {
|
||||
return protoimpl.X.MessageStringOf(x)
|
||||
}
|
||||
|
||||
func (*Account) ProtoMessage() {}
|
||||
|
||||
func (x *Account) ProtoReflect() protoreflect.Message {
|
||||
mi := &file_proxy_masque_config_proto_msgTypes[1]
|
||||
if x != nil {
|
||||
ms := protoimpl.X.MessageStateOf(protoimpl.Pointer(x))
|
||||
if ms.LoadMessageInfo() == nil {
|
||||
ms.StoreMessageInfo(mi)
|
||||
}
|
||||
return ms
|
||||
}
|
||||
return mi.MessageOf(x)
|
||||
}
|
||||
|
||||
// Deprecated: Use Account.ProtoReflect.Descriptor instead.
|
||||
func (*Account) Descriptor() ([]byte, []int) {
|
||||
return file_proxy_masque_config_proto_rawDescGZIP(), []int{1}
|
||||
}
|
||||
|
||||
func (x *Account) GetPassword() string {
|
||||
if x != nil {
|
||||
return x.Password
|
||||
}
|
||||
return ""
|
||||
}
|
||||
|
||||
type ServerConfig struct {
|
||||
state protoimpl.MessageState `protogen:"open.v1"`
|
||||
Users []*protocol.User `protobuf:"bytes,1,rep,name=users,proto3" json:"users,omitempty"`
|
||||
Address []string `protobuf:"bytes,2,rep,name=address,proto3" json:"address,omitempty"`
|
||||
Mtu uint32 `protobuf:"varint,3,opt,name=mtu,proto3" json:"mtu,omitempty"`
|
||||
unknownFields protoimpl.UnknownFields
|
||||
sizeCache protoimpl.SizeCache
|
||||
}
|
||||
|
||||
func (x *ServerConfig) Reset() {
|
||||
*x = ServerConfig{}
|
||||
mi := &file_proxy_masque_config_proto_msgTypes[2]
|
||||
ms := protoimpl.X.MessageStateOf(protoimpl.Pointer(x))
|
||||
ms.StoreMessageInfo(mi)
|
||||
}
|
||||
|
||||
func (x *ServerConfig) String() string {
|
||||
return protoimpl.X.MessageStringOf(x)
|
||||
}
|
||||
|
||||
func (*ServerConfig) ProtoMessage() {}
|
||||
|
||||
func (x *ServerConfig) ProtoReflect() protoreflect.Message {
|
||||
mi := &file_proxy_masque_config_proto_msgTypes[2]
|
||||
if x != nil {
|
||||
ms := protoimpl.X.MessageStateOf(protoimpl.Pointer(x))
|
||||
if ms.LoadMessageInfo() == nil {
|
||||
ms.StoreMessageInfo(mi)
|
||||
}
|
||||
return ms
|
||||
}
|
||||
return mi.MessageOf(x)
|
||||
}
|
||||
|
||||
// Deprecated: Use ServerConfig.ProtoReflect.Descriptor instead.
|
||||
func (*ServerConfig) Descriptor() ([]byte, []int) {
|
||||
return file_proxy_masque_config_proto_rawDescGZIP(), []int{2}
|
||||
}
|
||||
|
||||
func (x *ServerConfig) GetUsers() []*protocol.User {
|
||||
if x != nil {
|
||||
return x.Users
|
||||
}
|
||||
return nil
|
||||
}
|
||||
|
||||
func (x *ServerConfig) GetAddress() []string {
|
||||
if x != nil {
|
||||
return x.Address
|
||||
}
|
||||
return nil
|
||||
}
|
||||
|
||||
func (x *ServerConfig) GetMtu() uint32 {
|
||||
if x != nil {
|
||||
return x.Mtu
|
||||
}
|
||||
return 0
|
||||
}
|
||||
|
||||
var File_proxy_masque_config_proto protoreflect.FileDescriptor
|
||||
|
||||
const file_proxy_masque_config_proto_rawDesc = "" +
|
||||
"\n" +
|
||||
"\x19proxy/masque/config.proto\x12\x11xray.proxy.masque\x1a!common/protocol/server_spec.proto\"k\n" +
|
||||
"\x19proxy/masque/config.proto\x12\x11xray.proxy.masque\x1a!common/protocol/server_spec.proto\x1a\x1acommon/protocol/user.proto\"k\n" +
|
||||
"\fClientConfig\x12<\n" +
|
||||
"\x06server\x18\x01 \x01(\v2$.xray.common.protocol.ServerEndpointR\x06server\x12\x1d\n" +
|
||||
"\n" +
|
||||
"remote_dns\x18\x02 \x03(\tR\tremoteDnsBU\n" +
|
||||
"remote_dns\x18\x02 \x03(\tR\tremoteDns\"%\n" +
|
||||
"\aAccount\x12\x1a\n" +
|
||||
"\bpassword\x18\x01 \x01(\tR\bpassword\"l\n" +
|
||||
"\fServerConfig\x120\n" +
|
||||
"\x05users\x18\x01 \x03(\v2\x1a.xray.common.protocol.UserR\x05users\x12\x18\n" +
|
||||
"\aaddress\x18\x02 \x03(\tR\aaddress\x12\x10\n" +
|
||||
"\x03mtu\x18\x03 \x01(\rR\x03mtuBU\n" +
|
||||
"\x15com.xray.proxy.masqueP\x01Z&github.com/xtls/xray-core/proxy/masque\xaa\x02\x11Xray.Proxy.Masqueb\x06proto3"
|
||||
|
||||
var (
|
||||
@@ -97,18 +207,22 @@ func file_proxy_masque_config_proto_rawDescGZIP() []byte {
|
||||
return file_proxy_masque_config_proto_rawDescData
|
||||
}
|
||||
|
||||
var file_proxy_masque_config_proto_msgTypes = make([]protoimpl.MessageInfo, 1)
|
||||
var file_proxy_masque_config_proto_msgTypes = make([]protoimpl.MessageInfo, 3)
|
||||
var file_proxy_masque_config_proto_goTypes = []any{
|
||||
(*ClientConfig)(nil), // 0: xray.proxy.masque.ClientConfig
|
||||
(*protocol.ServerEndpoint)(nil), // 1: xray.common.protocol.ServerEndpoint
|
||||
(*Account)(nil), // 1: xray.proxy.masque.Account
|
||||
(*ServerConfig)(nil), // 2: xray.proxy.masque.ServerConfig
|
||||
(*protocol.ServerEndpoint)(nil), // 3: xray.common.protocol.ServerEndpoint
|
||||
(*protocol.User)(nil), // 4: xray.common.protocol.User
|
||||
}
|
||||
var file_proxy_masque_config_proto_depIdxs = []int32{
|
||||
1, // 0: xray.proxy.masque.ClientConfig.server:type_name -> xray.common.protocol.ServerEndpoint
|
||||
1, // [1:1] is the sub-list for method output_type
|
||||
1, // [1:1] is the sub-list for method input_type
|
||||
1, // [1:1] is the sub-list for extension type_name
|
||||
1, // [1:1] is the sub-list for extension extendee
|
||||
0, // [0:1] is the sub-list for field type_name
|
||||
3, // 0: xray.proxy.masque.ClientConfig.server:type_name -> xray.common.protocol.ServerEndpoint
|
||||
4, // 1: xray.proxy.masque.ServerConfig.users:type_name -> xray.common.protocol.User
|
||||
2, // [2:2] is the sub-list for method output_type
|
||||
2, // [2:2] is the sub-list for method input_type
|
||||
2, // [2:2] is the sub-list for extension type_name
|
||||
2, // [2:2] is the sub-list for extension extendee
|
||||
0, // [0:2] is the sub-list for field type_name
|
||||
}
|
||||
|
||||
func init() { file_proxy_masque_config_proto_init() }
|
||||
@@ -122,7 +236,7 @@ func file_proxy_masque_config_proto_init() {
|
||||
GoPackagePath: reflect.TypeOf(x{}).PkgPath(),
|
||||
RawDescriptor: unsafe.Slice(unsafe.StringData(file_proxy_masque_config_proto_rawDesc), len(file_proxy_masque_config_proto_rawDesc)),
|
||||
NumEnums: 0,
|
||||
NumMessages: 1,
|
||||
NumMessages: 3,
|
||||
NumExtensions: 0,
|
||||
NumServices: 0,
|
||||
},
|
||||
|
||||
@@ -7,8 +7,19 @@ option java_package = "com.xray.proxy.masque";
|
||||
option java_multiple_files = true;
|
||||
|
||||
import "common/protocol/server_spec.proto";
|
||||
import "common/protocol/user.proto";
|
||||
|
||||
message ClientConfig {
|
||||
xray.common.protocol.ServerEndpoint server = 1;
|
||||
repeated string remote_dns = 2;
|
||||
}
|
||||
|
||||
message Account {
|
||||
string password = 1;
|
||||
}
|
||||
|
||||
message ServerConfig {
|
||||
repeated xray.common.protocol.User users = 1;
|
||||
repeated string address = 2;
|
||||
uint32 mtu = 3;
|
||||
}
|
||||
|
||||
@@ -0,0 +1,80 @@
|
||||
package masque
|
||||
|
||||
import (
|
||||
"net/netip"
|
||||
"sync"
|
||||
|
||||
"github.com/xtls/xray-core/common/errors"
|
||||
)
|
||||
|
||||
type addressPool struct {
|
||||
mu sync.Mutex
|
||||
prefix netip.Prefix
|
||||
server netip.Addr
|
||||
first netip.Addr
|
||||
last netip.Addr
|
||||
next netip.Addr
|
||||
used map[netip.Addr]struct{}
|
||||
}
|
||||
|
||||
func newAddressPool(address netip.Prefix) (*addressPool, error) {
|
||||
server := address.Addr()
|
||||
if server.Is4In6() || server.Zone() != "" {
|
||||
return nil, errors.New("invalid address ", address)
|
||||
}
|
||||
prefix := address.Masked()
|
||||
last := lastAddr(prefix)
|
||||
if server == prefix.Addr() || server.Is4() && server == last {
|
||||
return nil, errors.New("address ", address, " is not a host address")
|
||||
}
|
||||
if server.Is4() {
|
||||
last = last.Prev()
|
||||
}
|
||||
first := prefix.Addr().Next()
|
||||
if first == last {
|
||||
return nil, errors.New("address ", address, " leaves no addresses to assign")
|
||||
}
|
||||
return &addressPool{
|
||||
prefix: prefix,
|
||||
server: server,
|
||||
first: first,
|
||||
last: last,
|
||||
next: first,
|
||||
used: make(map[netip.Addr]struct{}),
|
||||
}, nil
|
||||
}
|
||||
|
||||
func lastAddr(prefix netip.Prefix) netip.Addr {
|
||||
b := prefix.Addr().AsSlice()
|
||||
for i := prefix.Bits(); i < len(b)*8; i++ {
|
||||
b[i/8] |= 1 << (7 - i%8)
|
||||
}
|
||||
addr, _ := netip.AddrFromSlice(b)
|
||||
return addr
|
||||
}
|
||||
|
||||
func (p *addressPool) allocate() (netip.Addr, bool) {
|
||||
p.mu.Lock()
|
||||
defer p.mu.Unlock()
|
||||
for addr := p.next; ; {
|
||||
next := addr.Next()
|
||||
if addr == p.last {
|
||||
next = p.first
|
||||
}
|
||||
if _, found := p.used[addr]; !found && addr != p.server {
|
||||
p.used[addr] = struct{}{}
|
||||
p.next = next
|
||||
return addr, true
|
||||
}
|
||||
if next == p.next {
|
||||
return netip.Addr{}, false
|
||||
}
|
||||
addr = next
|
||||
}
|
||||
}
|
||||
|
||||
func (p *addressPool) release(addr netip.Addr) {
|
||||
p.mu.Lock()
|
||||
defer p.mu.Unlock()
|
||||
delete(p.used, addr)
|
||||
}
|
||||
@@ -0,0 +1,59 @@
|
||||
package masque
|
||||
|
||||
import (
|
||||
"net/netip"
|
||||
"testing"
|
||||
|
||||
"github.com/stretchr/testify/require"
|
||||
)
|
||||
|
||||
func allocateAll(p *addressPool) []netip.Addr {
|
||||
var addrs []netip.Addr
|
||||
for {
|
||||
addr, ok := p.allocate()
|
||||
if !ok {
|
||||
return addrs
|
||||
}
|
||||
addrs = append(addrs, addr)
|
||||
}
|
||||
}
|
||||
|
||||
func TestAddressPool(t *testing.T) {
|
||||
p, err := newAddressPool(netip.MustParsePrefix("10.0.0.1/29"))
|
||||
require.NoError(t, err)
|
||||
var want []netip.Addr
|
||||
for _, s := range []string{"10.0.0.2", "10.0.0.3", "10.0.0.4", "10.0.0.5", "10.0.0.6"} {
|
||||
want = append(want, netip.MustParseAddr(s))
|
||||
}
|
||||
require.Equal(t, want, allocateAll(p))
|
||||
|
||||
p.release(netip.MustParseAddr("10.0.0.4"))
|
||||
addr, ok := p.allocate()
|
||||
require.True(t, ok)
|
||||
require.Equal(t, netip.MustParseAddr("10.0.0.4"), addr)
|
||||
_, ok = p.allocate()
|
||||
require.False(t, ok)
|
||||
|
||||
p, err = newAddressPool(netip.MustParsePrefix("fd00::1/126"))
|
||||
require.NoError(t, err)
|
||||
require.Equal(t, []netip.Addr{netip.MustParseAddr("fd00::2"), netip.MustParseAddr("fd00::3")}, allocateAll(p))
|
||||
|
||||
p, err = newAddressPool(netip.MustParsePrefix("10.0.0.2/30"))
|
||||
require.NoError(t, err)
|
||||
require.Equal(t, []netip.Addr{netip.MustParseAddr("10.0.0.1")}, allocateAll(p))
|
||||
}
|
||||
|
||||
func TestAddressPoolRejects(t *testing.T) {
|
||||
for _, s := range []string{
|
||||
"10.0.0.0/24",
|
||||
"10.0.0.255/24",
|
||||
"10.0.0.1/31",
|
||||
"10.0.0.1/32",
|
||||
"fd00::1/127",
|
||||
"fd00::1/128",
|
||||
"::ffff:10.0.0.1/120",
|
||||
} {
|
||||
_, err := newAddressPool(netip.MustParsePrefix(s))
|
||||
require.Error(t, err, s)
|
||||
}
|
||||
}
|
||||
@@ -0,0 +1,550 @@
|
||||
package masque
|
||||
|
||||
import (
|
||||
"context"
|
||||
go_errors "errors"
|
||||
"io"
|
||||
stdnet "net"
|
||||
"net/http"
|
||||
"net/netip"
|
||||
"slices"
|
||||
"sync"
|
||||
|
||||
"golang.zx2c4.com/wireguard/tun"
|
||||
|
||||
"github.com/xtls/xray-core/common"
|
||||
"github.com/xtls/xray-core/common/buf"
|
||||
c "github.com/xtls/xray-core/common/ctx"
|
||||
"github.com/xtls/xray-core/common/errors"
|
||||
"github.com/xtls/xray-core/common/log"
|
||||
"github.com/xtls/xray-core/common/net"
|
||||
"github.com/xtls/xray-core/common/protocol"
|
||||
"github.com/xtls/xray-core/common/session"
|
||||
"github.com/xtls/xray-core/core"
|
||||
"github.com/xtls/xray-core/features/routing"
|
||||
"github.com/xtls/xray-core/proxy/wireguard"
|
||||
"github.com/xtls/xray-core/transport"
|
||||
"github.com/xtls/xray-core/transport/internet"
|
||||
"github.com/xtls/xray-core/transport/internet/masque"
|
||||
"github.com/xtls/xray-core/transport/internet/masque/connectip"
|
||||
"github.com/xtls/xray-core/transport/internet/stat"
|
||||
"github.com/xtls/xray-core/transport/internet/tls"
|
||||
)
|
||||
|
||||
const (
|
||||
authenticateHeader = `Basic realm="masque", charset="UTF-8"`
|
||||
tunnelQueueSize = 512
|
||||
)
|
||||
|
||||
type Server struct {
|
||||
validator *validator
|
||||
dispatcher routing.Dispatcher
|
||||
ctx context.Context
|
||||
tag string
|
||||
sniffing session.SniffingRequest
|
||||
mtu int
|
||||
|
||||
dev tun.Device
|
||||
pools []*addressPool
|
||||
local []netip.Addr
|
||||
|
||||
mu sync.RWMutex
|
||||
tunnels map[netip.Addr]*serverTunnel
|
||||
closed bool
|
||||
started bool
|
||||
}
|
||||
|
||||
type serverTunnel struct {
|
||||
conn stat.Connection
|
||||
ipConn *connectip.Conn
|
||||
user *protocol.MemoryUser
|
||||
addrs []netip.Addr
|
||||
queue chan *buf.Buffer
|
||||
done chan struct{}
|
||||
|
||||
mu sync.Mutex
|
||||
conns map[net.Conn]struct{}
|
||||
}
|
||||
|
||||
func newServerTunnel(conn stat.Connection, user *protocol.MemoryUser) *serverTunnel {
|
||||
return &serverTunnel{
|
||||
conn: conn,
|
||||
user: user,
|
||||
queue: make(chan *buf.Buffer, tunnelQueueSize),
|
||||
done: make(chan struct{}),
|
||||
conns: make(map[net.Conn]struct{}),
|
||||
}
|
||||
}
|
||||
|
||||
func (t *serverTunnel) send(b *buf.Buffer) bool {
|
||||
select {
|
||||
case <-t.done:
|
||||
return false
|
||||
default:
|
||||
}
|
||||
select {
|
||||
case t.queue <- b:
|
||||
return true
|
||||
default:
|
||||
return false
|
||||
}
|
||||
}
|
||||
|
||||
func (t *serverTunnel) track(conn net.Conn) bool {
|
||||
t.mu.Lock()
|
||||
defer t.mu.Unlock()
|
||||
if t.conns == nil {
|
||||
return false
|
||||
}
|
||||
t.conns[conn] = struct{}{}
|
||||
return true
|
||||
}
|
||||
|
||||
func (t *serverTunnel) untrack(conn net.Conn) {
|
||||
t.mu.Lock()
|
||||
delete(t.conns, conn)
|
||||
t.mu.Unlock()
|
||||
}
|
||||
|
||||
func (t *serverTunnel) close() {
|
||||
t.mu.Lock()
|
||||
conns := t.conns
|
||||
if conns != nil {
|
||||
t.conns = nil
|
||||
close(t.done)
|
||||
}
|
||||
t.mu.Unlock()
|
||||
for conn := range conns {
|
||||
conn.Close()
|
||||
}
|
||||
}
|
||||
|
||||
func NewServer(ctx context.Context, config *ServerConfig) (*Server, error) {
|
||||
v := core.MustFromContext(ctx)
|
||||
|
||||
streamSettings := session.StreamSettingsFromContext(ctx).(*internet.MemoryStreamConfig)
|
||||
if _, ok := streamSettings.ProtocolSettings.(*masque.Config); !ok {
|
||||
return nil, errors.New("not masque transport")
|
||||
}
|
||||
if tls.ConfigFromStreamSettings(streamSettings) == nil {
|
||||
return nil, errors.New(`MASQUE requires "security": "tls"`)
|
||||
}
|
||||
|
||||
users := newValidator()
|
||||
for _, user := range config.Users {
|
||||
u, err := user.ToMemoryUser()
|
||||
if err != nil {
|
||||
return nil, errors.New("failed to get MASQUE user").Base(err)
|
||||
}
|
||||
if err := users.add(u); err != nil {
|
||||
return nil, errors.New("failed to add user").Base(err)
|
||||
}
|
||||
}
|
||||
|
||||
var pools []*addressPool
|
||||
var local []netip.Addr
|
||||
for _, s := range config.Address {
|
||||
prefix, err := netip.ParsePrefix(s)
|
||||
if err != nil {
|
||||
return nil, errors.New("invalid address ", s).Base(err)
|
||||
}
|
||||
if slices.ContainsFunc(local, func(addr netip.Addr) bool { return addr.Is4() == prefix.Addr().Is4() }) {
|
||||
return nil, errors.New("only one address per IP family is supported")
|
||||
}
|
||||
pool, err := newAddressPool(prefix)
|
||||
if err != nil {
|
||||
return nil, err
|
||||
}
|
||||
pools = append(pools, pool)
|
||||
local = append(local, prefix.Addr())
|
||||
}
|
||||
if len(pools) == 0 {
|
||||
return nil, errors.New("no address to assign")
|
||||
}
|
||||
|
||||
mtu := int(config.Mtu)
|
||||
if mtu == 0 {
|
||||
mtu = masque.MinPacketSize
|
||||
}
|
||||
dev, _, gstack, err := wireguard.CreateNetTUN(local, nil, mtu, false)
|
||||
if err != nil {
|
||||
return nil, err
|
||||
}
|
||||
|
||||
s := &Server{
|
||||
validator: users,
|
||||
dispatcher: v.GetFeature(routing.DispatcherType()).(routing.Dispatcher),
|
||||
ctx: core.ToBackgroundDetachedContext(ctx),
|
||||
mtu: mtu,
|
||||
dev: dev,
|
||||
pools: pools,
|
||||
local: local,
|
||||
tunnels: make(map[netip.Addr]*serverTunnel),
|
||||
}
|
||||
if inbound := session.InboundFromContext(ctx); inbound != nil {
|
||||
s.tag = inbound.Tag
|
||||
}
|
||||
if content := session.ContentFromContext(ctx); content != nil {
|
||||
s.sniffing = content.SniffingRequest
|
||||
}
|
||||
wireguard.CreateForwarder(gstack, s.handleConnection)
|
||||
return s, nil
|
||||
}
|
||||
|
||||
func (s *Server) Start() error {
|
||||
s.mu.Lock()
|
||||
defer s.mu.Unlock()
|
||||
if s.started || s.closed {
|
||||
return nil
|
||||
}
|
||||
s.started = true
|
||||
go s.readFromStack()
|
||||
return nil
|
||||
}
|
||||
|
||||
func (s *Server) Close() error {
|
||||
s.mu.Lock()
|
||||
if s.closed {
|
||||
s.mu.Unlock()
|
||||
return nil
|
||||
}
|
||||
s.closed = true
|
||||
var tunnels []*serverTunnel
|
||||
for _, t := range s.tunnels {
|
||||
if !slices.Contains(tunnels, t) {
|
||||
tunnels = append(tunnels, t)
|
||||
}
|
||||
}
|
||||
s.mu.Unlock()
|
||||
for _, t := range tunnels {
|
||||
t.conn.Close()
|
||||
}
|
||||
return s.dev.Close()
|
||||
}
|
||||
|
||||
func (s *Server) AddUser(ctx context.Context, user *protocol.MemoryUser) error {
|
||||
return s.validator.add(user)
|
||||
}
|
||||
|
||||
func (s *Server) RemoveUser(ctx context.Context, email string) error {
|
||||
user, err := s.validator.delByEmail(email)
|
||||
if err != nil {
|
||||
return err
|
||||
}
|
||||
s.mu.RLock()
|
||||
var conns []stat.Connection
|
||||
for _, t := range s.tunnels {
|
||||
if t.user == user && !slices.Contains(conns, t.conn) {
|
||||
conns = append(conns, t.conn)
|
||||
}
|
||||
}
|
||||
s.mu.RUnlock()
|
||||
for _, conn := range conns {
|
||||
conn.Close()
|
||||
}
|
||||
return nil
|
||||
}
|
||||
|
||||
func (s *Server) GetUser(ctx context.Context, email string) *protocol.MemoryUser {
|
||||
return s.validator.getByEmail(email)
|
||||
}
|
||||
|
||||
func (s *Server) GetUsers(ctx context.Context) []*protocol.MemoryUser {
|
||||
return s.validator.getAll()
|
||||
}
|
||||
|
||||
func (s *Server) GetUsersCount(context.Context) int64 {
|
||||
return s.validator.count()
|
||||
}
|
||||
|
||||
func (s *Server) Network() []net.Network {
|
||||
return []net.Network{net.Network_TCP}
|
||||
}
|
||||
|
||||
func (s *Server) Process(ctx context.Context, network net.Network, conn stat.Connection, dispatcher routing.Dispatcher) error {
|
||||
sconn, ok := stat.TryUnwrapStatsConn(conn).(*masque.ServerConn)
|
||||
if !ok {
|
||||
return errors.New("not a MASQUE connection")
|
||||
}
|
||||
inbound := session.InboundFromContext(ctx)
|
||||
inbound.Name = "masque"
|
||||
inbound.CanSpliceCopy = 3
|
||||
|
||||
name, pass, _ := sconn.Request().BasicAuth()
|
||||
user := s.validator.get(name, pass)
|
||||
if user == nil {
|
||||
sconn.Reject(http.StatusUnauthorized, http.Header{"WWW-Authenticate": {authenticateHeader}})
|
||||
log.Record(&log.AccessMessage{
|
||||
From: conn.RemoteAddr(),
|
||||
To: "",
|
||||
Status: log.AccessRejected,
|
||||
Reason: errors.New("invalid credentials"),
|
||||
})
|
||||
return errors.New("MASQUE: authentication failed for ", name)
|
||||
}
|
||||
inbound.User = user
|
||||
|
||||
t := newServerTunnel(conn, user)
|
||||
for _, pool := range s.pools {
|
||||
if addr, ok := pool.allocate(); ok {
|
||||
t.addrs = append(t.addrs, addr)
|
||||
}
|
||||
}
|
||||
defer s.release(t)
|
||||
if len(t.addrs) == 0 {
|
||||
sconn.Reject(http.StatusServiceUnavailable, nil)
|
||||
return errors.New("MASQUE: no address left to assign")
|
||||
}
|
||||
|
||||
ipConn, err := sconn.Accept()
|
||||
if err != nil {
|
||||
return errors.New("MASQUE: failed to accept the tunnel").Base(err)
|
||||
}
|
||||
t.ipConn = ipConn
|
||||
if !s.register(t) {
|
||||
return errors.New("MASQUE: server closed")
|
||||
}
|
||||
if !s.validator.contains(user) {
|
||||
return errors.New("MASQUE: user ", name, " was removed")
|
||||
}
|
||||
go s.writeToTunnel(t)
|
||||
|
||||
prefixes := make([]netip.Prefix, len(t.addrs))
|
||||
for i, addr := range t.addrs {
|
||||
prefixes[i] = netip.PrefixFrom(addr, addr.BitLen())
|
||||
}
|
||||
if err := ipConn.AssignAddresses(prefixes); err != nil {
|
||||
return err
|
||||
}
|
||||
if err := ipConn.AdvertiseRoute(fullRoutes(t.addrs)); err != nil {
|
||||
return err
|
||||
}
|
||||
go serveAddressRequests(t)
|
||||
|
||||
ctx = log.ContextWithAccessMessage(ctx, &log.AccessMessage{
|
||||
From: conn.RemoteAddr(),
|
||||
To: "",
|
||||
Status: log.AccessAccepted,
|
||||
Email: user.Email,
|
||||
})
|
||||
errors.LogInfo(ctx, "MASQUE: tunnel from ", inbound.Source, " assigned ", t.addrs)
|
||||
return s.readFromTunnel(t)
|
||||
}
|
||||
|
||||
func fullRoutes(addrs []netip.Addr) []connectip.IPRoute {
|
||||
var routes []connectip.IPRoute
|
||||
if slices.ContainsFunc(addrs, netip.Addr.Is4) {
|
||||
routes = append(routes, connectip.IPRoute{StartIP: netip.IPv4Unspecified(), EndIP: netip.AddrFrom4([4]byte{255, 255, 255, 255})})
|
||||
}
|
||||
if slices.ContainsFunc(addrs, netip.Addr.Is6) {
|
||||
routes = append(routes, connectip.IPRoute{StartIP: netip.IPv6Unspecified(), EndIP: netip.AddrFrom16([16]byte{0: 0xff, 1: 0xff, 2: 0xff, 3: 0xff, 4: 0xff, 5: 0xff, 6: 0xff, 7: 0xff, 8: 0xff, 9: 0xff, 10: 0xff, 11: 0xff, 12: 0xff, 13: 0xff, 14: 0xff, 15: 0xff})})
|
||||
}
|
||||
return routes
|
||||
}
|
||||
|
||||
func serveAddressRequests(t *serverTunnel) {
|
||||
for {
|
||||
req, err := t.ipConn.ReceiveAddressRequest(context.Background())
|
||||
if err != nil {
|
||||
return
|
||||
}
|
||||
assigned := make([]netip.Prefix, len(req.Prefixes))
|
||||
used := make(map[netip.Addr]bool)
|
||||
for i, requested := range req.Prefixes {
|
||||
for _, addr := range t.addrs {
|
||||
if addr.Is4() == requested.Addr().Is4() && !used[addr] {
|
||||
used[addr] = true
|
||||
assigned[i] = netip.PrefixFrom(addr, addr.BitLen())
|
||||
break
|
||||
}
|
||||
}
|
||||
}
|
||||
var additional []netip.Prefix
|
||||
for _, addr := range t.addrs {
|
||||
if !used[addr] {
|
||||
additional = append(additional, netip.PrefixFrom(addr, addr.BitLen()))
|
||||
}
|
||||
}
|
||||
if err := req.Respond(assigned, additional); err != nil {
|
||||
return
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
func (s *Server) register(t *serverTunnel) bool {
|
||||
s.mu.Lock()
|
||||
defer s.mu.Unlock()
|
||||
if s.closed {
|
||||
return false
|
||||
}
|
||||
for _, addr := range t.addrs {
|
||||
s.tunnels[addr] = t
|
||||
}
|
||||
return true
|
||||
}
|
||||
|
||||
func (s *Server) release(t *serverTunnel) {
|
||||
s.mu.Lock()
|
||||
for _, addr := range t.addrs {
|
||||
if s.tunnels[addr] == t {
|
||||
delete(s.tunnels, addr)
|
||||
}
|
||||
}
|
||||
s.mu.Unlock()
|
||||
t.close()
|
||||
for _, addr := range t.addrs {
|
||||
for _, pool := range s.pools {
|
||||
if pool.prefix.Contains(addr) {
|
||||
pool.release(addr)
|
||||
}
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
func (s *Server) lookup(addr netip.Addr) *serverTunnel {
|
||||
s.mu.RLock()
|
||||
defer s.mu.RUnlock()
|
||||
return s.tunnels[addr]
|
||||
}
|
||||
|
||||
func (s *Server) inPool(addr netip.Addr) bool {
|
||||
return slices.ContainsFunc(s.pools, func(pool *addressPool) bool { return pool.prefix.Contains(addr) })
|
||||
}
|
||||
|
||||
func (s *Server) readFromTunnel(t *serverTunnel) error {
|
||||
b := make([]byte, 1<<16)
|
||||
for {
|
||||
n, err := t.conn.Read(b)
|
||||
if err != nil {
|
||||
if go_errors.Is(err, io.ErrShortBuffer) {
|
||||
continue
|
||||
}
|
||||
if go_errors.Is(err, stdnet.ErrClosed) || go_errors.Is(err, io.EOF) {
|
||||
return nil
|
||||
}
|
||||
return err
|
||||
}
|
||||
dst, ok := packetDestination(b[:n])
|
||||
if !ok || dst.IsLinkLocalUnicast() || dst.IsMulticast() {
|
||||
continue
|
||||
}
|
||||
if other := s.lookup(dst); other != nil {
|
||||
if other != t {
|
||||
packet := buf.NewWithSize(int32(n))
|
||||
packet.Write(b[:n])
|
||||
if !other.send(packet) {
|
||||
packet.Release()
|
||||
}
|
||||
}
|
||||
continue
|
||||
}
|
||||
if s.inPool(dst) && !slices.Contains(s.local, dst) {
|
||||
continue
|
||||
}
|
||||
s.dev.Write([][]byte{b[:n]}, 0)
|
||||
}
|
||||
}
|
||||
|
||||
func (s *Server) readFromStack() {
|
||||
sizes := []int{0}
|
||||
var b *buf.Buffer
|
||||
for {
|
||||
if b == nil {
|
||||
b = buf.NewWithSize(int32(s.mtu))
|
||||
}
|
||||
b.Clear()
|
||||
if _, err := s.dev.Read([][]byte{b.Extend(int32(s.mtu))}, sizes, 0); err != nil {
|
||||
b.Release()
|
||||
return
|
||||
}
|
||||
b.Resize(0, int32(sizes[0]))
|
||||
dst, ok := packetDestination(b.Bytes())
|
||||
if !ok {
|
||||
continue
|
||||
}
|
||||
if t := s.lookup(dst); t != nil && t.send(b) {
|
||||
b = nil
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
func (s *Server) writeToTunnel(t *serverTunnel) {
|
||||
for {
|
||||
select {
|
||||
case b := <-t.queue:
|
||||
_, err := t.conn.Write(b.Bytes())
|
||||
b.Release()
|
||||
if ptb, ok := go_errors.AsType[*masque.PacketTooBigError](err); ok {
|
||||
s.dev.Write([][]byte{ptb.ICMP}, 0)
|
||||
}
|
||||
case <-t.done:
|
||||
return
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
func packetDestination(packet []byte) (netip.Addr, bool) {
|
||||
if len(packet) == 0 {
|
||||
return netip.Addr{}, false
|
||||
}
|
||||
switch packet[0] >> 4 {
|
||||
case 4:
|
||||
if len(packet) >= 20 {
|
||||
return netip.AddrFrom4([4]byte(packet[16:20])), true
|
||||
}
|
||||
case 6:
|
||||
if len(packet) >= 40 {
|
||||
return netip.AddrFrom16([16]byte(packet[24:40])), true
|
||||
}
|
||||
}
|
||||
return netip.Addr{}, false
|
||||
}
|
||||
|
||||
func (s *Server) handleConnection(conn net.Conn, dest net.Destination) {
|
||||
defer conn.Close()
|
||||
source := net.DestinationFromAddr(conn.RemoteAddr())
|
||||
addr, _ := netip.AddrFromSlice(source.Address.IP())
|
||||
t := s.lookup(addr.Unmap())
|
||||
if t == nil || !t.track(conn) {
|
||||
errors.LogInfo(s.ctx, "MASQUE: no tunnel for ", source, " to ", dest)
|
||||
return
|
||||
}
|
||||
defer t.untrack(conn)
|
||||
|
||||
ctx, cancel := context.WithCancel(s.ctx)
|
||||
defer cancel()
|
||||
ctx = c.ContextWithID(ctx, session.NewID())
|
||||
inbound := session.Inbound{
|
||||
Name: "masque",
|
||||
Tag: s.tag,
|
||||
CanSpliceCopy: 3,
|
||||
Source: source,
|
||||
User: t.user,
|
||||
}
|
||||
ctx = session.ContextWithInbound(ctx, &inbound)
|
||||
ctx = session.ContextWithContent(ctx, &session.Content{
|
||||
SniffingRequest: s.sniffing,
|
||||
})
|
||||
ctx = session.SubContextFromMuxInbound(ctx)
|
||||
ctx = log.ContextWithAccessMessage(ctx, &log.AccessMessage{
|
||||
From: source,
|
||||
To: dest,
|
||||
Status: log.AccessAccepted,
|
||||
Email: t.user.Email,
|
||||
})
|
||||
errors.LogInfo(ctx, "processing from ", source, " to ", dest)
|
||||
|
||||
link := &transport.Link{
|
||||
Reader: &buf.TimeoutWrapperReader{Reader: buf.NewReader(conn)},
|
||||
Writer: buf.NewWriter(conn),
|
||||
}
|
||||
if err := s.dispatcher.DispatchLink(ctx, dest, link); err != nil {
|
||||
errors.LogError(ctx, errors.New("connection closed").Base(err))
|
||||
}
|
||||
}
|
||||
|
||||
func init() {
|
||||
common.Must(common.RegisterConfig((*ServerConfig)(nil), func(ctx context.Context, config interface{}) (interface{}, error) {
|
||||
return NewServer(ctx, config.(*ServerConfig))
|
||||
}))
|
||||
}
|
||||
@@ -0,0 +1,328 @@
|
||||
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)
|
||||
}
|
||||
}
|
||||
Reference in New Issue
Block a user