diff --git a/infra/conf/masque.go b/infra/conf/masque.go index cb9c201a2..b128ee593 100644 --- a/infra/conf/masque.go +++ b/infra/conf/masque.go @@ -2,9 +2,11 @@ package conf import ( "net/netip" + "strings" "github.com/xtls/xray-core/common/errors" "github.com/xtls/xray-core/common/protocol" + "github.com/xtls/xray-core/common/serial" "github.com/xtls/xray-core/proxy/masque" "google.golang.org/protobuf/proto" ) @@ -35,3 +37,67 @@ func (c *MasqueClientConfig) Build() (proto.Message, error) { RemoteDns: c.RemoteDNS, }, nil } + +type MasqueUserConfig struct { + Pass string `json:"pass"` + Level uint32 `json:"level"` + Email string `json:"email"` +} + +type MasqueServerConfig struct { + Users []*MasqueUserConfig `json:"users"` + Clients []*MasqueUserConfig `json:"clients"` + Address []string `json:"address"` + MTU uint32 `json:"mtu"` +} + +func (c *MasqueServerConfig) Build() (proto.Message, error) { + if c.Clients != nil { + c.Users = c.Clients + } + config := &masque.ServerConfig{ + Address: c.Address, + Mtu: c.MTU, + } + emails := make(map[string]bool) + for _, user := range c.Users { + if user.Email == "" { + return nil, errors.New(`MASQUE: "email" is empty`) + } + if strings.Contains(user.Email, ":") { + return nil, errors.New(`MASQUE: invalid "email" `, user.Email) + } + if user.Pass == "" { + return nil, errors.New(`MASQUE: "pass" of `, user.Email, ` is empty`) + } + email := strings.ToLower(user.Email) + if emails[email] { + return nil, errors.New(`MASQUE: duplicate "email" `, user.Email) + } + emails[email] = true + config.Users = append(config.Users, &protocol.User{ + Email: user.Email, + Level: user.Level, + Account: serial.ToTypedMessage(&masque.Account{Password: user.Pass}), + }) + } + if len(c.Address) == 0 { + return nil, errors.New(`MASQUE: "address" is not set`) + } + var v4, v6 bool + for _, s := range c.Address { + prefix, err := netip.ParsePrefix(s) + if err != nil { + return nil, errors.New(`MASQUE: invalid "address" `, s).Base(err) + } + if prefix.Addr().Is4() && v4 || prefix.Addr().Is6() && v6 { + return nil, errors.New(`MASQUE: "address" takes at most one IPv4 and one IPv6 prefix`) + } + v4 = v4 || prefix.Addr().Is4() + v6 = v6 || prefix.Addr().Is6() + } + if c.MTU != 0 && (c.MTU < 1280 || c.MTU > 65535) { + return nil, errors.New(`MASQUE: "mtu" must be between 1280 and 65535`) + } + return config, nil +} diff --git a/infra/conf/masque_test.go b/infra/conf/masque_test.go index 806bb4842..14ec6eedf 100644 --- a/infra/conf/masque_test.go +++ b/infra/conf/masque_test.go @@ -4,7 +4,10 @@ import ( "encoding/json" "testing" + "github.com/xtls/xray-core/common/protocol" + "github.com/xtls/xray-core/common/serial" . "github.com/xtls/xray-core/infra/conf" + masqueproxy "github.com/xtls/xray-core/proxy/masque" "github.com/xtls/xray-core/transport/internet/masque" ) @@ -37,6 +40,14 @@ func TestMasqueConfig(t *testing.T) { Parser: loadJSON(creator), Output: &masque.Config{Path: "/masque/ip?target=*&ipproto=*"}, }, + { + Input: `{"user": "u", "pass": "p:q", "headers": {"X-Token": "a"}}`, + Parser: loadJSON(creator), + Output: &masque.Config{ + Path: "/.well-known/masque/ip/*/*/", + Headers: map[string]string{"Authorization": "Basic dTpwOnE=", "X-Token": "a"}, + }, + }, }) for _, input := range []string{ @@ -47,6 +58,8 @@ func TestMasqueConfig(t *testing.T) { `{"headers": {"Capsule-Protocol": "?0"}}`, `{"headers": {"X Token": "a"}}`, `{"headers": {"X-Token": "a\r\nb"}}`, + `{"user": "u:v", "pass": "p"}`, + `{"user": "u", "pass": "p", "headers": {"authorization": "Basic dTpw"}}`, } { if _, err := loadJSON(creator)(input); err == nil { t.Errorf("expected an error for %s", input) @@ -83,3 +96,93 @@ func TestMasqueOutboundConfig(t *testing.T) { } } } + +func TestMasqueServerConfig(t *testing.T) { + creator := func() Buildable { + return new(MasqueServerConfig) + } + + runMultiTestCase(t, []TestCase{ + { + Input: `{ + "users": [{"email": "u@example.com", "pass": "p", "level": 1}], + "address": ["10.13.0.1/24", "fd13::1/64"], + "mtu": 1400 + }`, + Parser: loadJSON(creator), + Output: &masqueproxy.ServerConfig{ + Users: []*protocol.User{{ + Email: "u@example.com", + Level: 1, + Account: serial.ToTypedMessage(&masqueproxy.Account{Password: "p"}), + }}, + Address: []string{"10.13.0.1/24", "fd13::1/64"}, + Mtu: 1400, + }, + }, + { + Input: `{"clients": [{"email": "u", "pass": "p:q"}], "address": ["10.13.0.1/24"]}`, + Parser: loadJSON(creator), + Output: &masqueproxy.ServerConfig{ + Users: []*protocol.User{{ + Email: "u", + Account: serial.ToTypedMessage(&masqueproxy.Account{Password: "p:q"}), + }}, + Address: []string{"10.13.0.1/24"}, + }, + }, + { + Input: `{"address": ["10.13.0.1/24"]}`, + Parser: loadJSON(creator), + Output: &masqueproxy.ServerConfig{ + Address: []string{"10.13.0.1/24"}, + }, + }, + }) + + for _, input := range []string{ + `{"users": [{"email": "u:v", "pass": "p"}], "address": ["10.13.0.1/24"]}`, + `{"users": [{"email": "", "pass": "p"}], "address": ["10.13.0.1/24"]}`, + `{"users": [{"pass": "p"}], "address": ["10.13.0.1/24"]}`, + `{"users": [{"email": "u", "pass": ""}], "address": ["10.13.0.1/24"]}`, + `{"users": [{"email": "u", "pass": "p"}, {"email": "U", "pass": "q"}], "address": ["10.13.0.1/24"]}`, + `{"users": [{"email": "u", "pass": "p"}]}`, + `{"users": [{"email": "u", "pass": "p"}], "address": ["10.13.0.1"]}`, + `{"users": [{"email": "u", "pass": "p"}], "address": ["10.13.0.1/24", "10.14.0.1/24"]}`, + `{"users": [{"email": "u", "pass": "p"}], "address": ["fd13::1/64", "fd14::1/64"]}`, + `{"users": [{"email": "u", "pass": "p"}], "address": ["10.13.0.1/24"], "mtu": 1000}`, + `{"users": [{"email": "u", "pass": "p"}], "address": ["10.13.0.1/24"], "mtu": 70000}`, + } { + if _, err := loadJSON(creator)(input); err == nil { + t.Errorf("expected an error for %s", input) + } + } +} + +func TestMasqueInboundConfig(t *testing.T) { + build := func(s string) error { + c := new(InboundDetourConfig) + if err := json.Unmarshal([]byte(s), c); err != nil { + return err + } + _, err := c.Build() + return err + } + + if err := build(`{ + "protocol": "masque", + "port": 443, + "settings": {"users": [{"email": "u@example.com", "pass": "p"}], "address": ["10.13.0.1/24"]}, + "streamSettings": {"network": "masque", "security": "tls"} + }`); err != nil { + t.Error(err) + } + if err := build(`{ + "protocol": "vless", + "port": 443, + "settings": {"users": [{"id": "27848739-7e62-4138-9fd3-098a63964b6b"}], "decryption": "none"}, + "streamSettings": {"network": "masque", "security": "tls"} + }`); err == nil { + t.Error("expected an error for the masque transport on a vless inbound") + } +} diff --git a/infra/conf/transport_method.go b/infra/conf/transport_method.go index e6ba86f5c..eadc0519d 100644 --- a/infra/conf/transport_method.go +++ b/infra/conf/transport_method.go @@ -1,7 +1,9 @@ package conf import ( + "encoding/base64" "encoding/json" + "maps" "math/big" "net/url" "sort" @@ -791,6 +793,8 @@ func (c *HysteriaConfig) Build() (proto.Message, error) { type MasqueConfig struct { Host string `json:"host"` Path string `json:"path"` + User string `json:"user"` + Pass string `json:"pass"` Headers map[string]string `json:"headers"` } @@ -819,12 +823,27 @@ func (c *MasqueConfig) Build() (proto.Message, error) { switch strings.ToLower(k) { case "host", "capsule-protocol": return nil, errors.New(`"headers" can't contain "`, k, `"`) + case "authorization": + if c.User != "" || c.Pass != "" { + return nil, errors.New(`"headers" can't contain "`, k, `" when "user" or "pass" is set`) + } } } + headers := c.Headers + if c.User != "" || c.Pass != "" { + if strings.Contains(c.User, ":") { + return nil, errors.New(`invalid "user": `, c.User) + } + headers = maps.Clone(c.Headers) + if headers == nil { + headers = make(map[string]string) + } + headers["Authorization"] = "Basic " + base64.StdEncoding.EncodeToString([]byte(c.User+":"+c.Pass)) + } return &masque.Config{ Host: c.Host, Path: path, - Headers: c.Headers, + Headers: headers, }, nil } diff --git a/infra/conf/xray.go b/infra/conf/xray.go index 46d3ae2e4..1ce1bd06d 100644 --- a/infra/conf/xray.go +++ b/infra/conf/xray.go @@ -33,6 +33,7 @@ var ( "trojan": func() interface{} { return new(TrojanServerConfig) }, "wireguard": func() interface{} { return &WireGuardConfig{IsClient: false} }, "hysteria": func() interface{} { return new(HysteriaServerConfig) }, + "masque": func() interface{} { return new(MasqueServerConfig) }, "tun": func() interface{} { return new(TunConfig) }, }, "protocol", "settings") @@ -205,6 +206,9 @@ func (c *InboundDetourConfig) Build() (*core.InboundHandlerConfig, error) { if err != nil { return nil, errors.New("failed to build inbound handler for protocol ", c.Protocol).Base(err) } + if _, ok := ts.(*masque.ServerConfig); !ok && receiverSettings.StreamSettings != nil && receiverSettings.StreamSettings.ProtocolName == "masque" { + return nil, errors.New("the masque transport can only be used by the masque inbound") + } return &core.InboundHandlerConfig{ Tag: c.Tag, diff --git a/main/commands/all/api/inbound_user_add.go b/main/commands/all/api/inbound_user_add.go index 81afc0d29..1b150bb5a 100644 --- a/main/commands/all/api/inbound_user_add.go +++ b/main/commands/all/api/inbound_user_add.go @@ -12,6 +12,7 @@ import ( "github.com/xtls/xray-core/core" "github.com/xtls/xray-core/infra/conf" "github.com/xtls/xray-core/infra/conf/serial" + "github.com/xtls/xray-core/proxy/masque" "github.com/xtls/xray-core/proxy/shadowsocks" "github.com/xtls/xray-core/proxy/shadowsocks_2022" "github.com/xtls/xray-core/proxy/trojan" @@ -88,6 +89,8 @@ func extractInboundUsers(inb *core.InboundHandlerConfig) []*protocol.User { return ty.Users case *shadowsocks_2022.MultiUserServerConfig: return ty.Users + case *masque.ServerConfig: + return ty.Users default: fmt.Println("unsupported inbound type") } diff --git a/proxy/masque/account.go b/proxy/masque/account.go new file mode 100644 index 000000000..9b505bcf1 --- /dev/null +++ b/proxy/masque/account.go @@ -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)) +} diff --git a/proxy/masque/account_test.go b/proxy/masque/account_test.go new file mode 100644 index 000000000..907faa9f6 --- /dev/null +++ b/proxy/masque/account_test.go @@ -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()) +} diff --git a/proxy/masque/config.pb.go b/proxy/masque/config.pb.go index 2d5d8e69c..bcb2d7c1f 100644 --- a/proxy/masque/config.pb.go +++ b/proxy/masque/config.pb.go @@ -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, }, diff --git a/proxy/masque/config.proto b/proxy/masque/config.proto index 0b402908a..3af0879dd 100644 --- a/proxy/masque/config.proto +++ b/proxy/masque/config.proto @@ -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; +} diff --git a/proxy/masque/pool.go b/proxy/masque/pool.go new file mode 100644 index 000000000..3f24229d0 --- /dev/null +++ b/proxy/masque/pool.go @@ -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) +} diff --git a/proxy/masque/pool_test.go b/proxy/masque/pool_test.go new file mode 100644 index 000000000..f01421ba7 --- /dev/null +++ b/proxy/masque/pool_test.go @@ -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) + } +} diff --git a/proxy/masque/server.go b/proxy/masque/server.go new file mode 100644 index 000000000..f9cc2ff37 --- /dev/null +++ b/proxy/masque/server.go @@ -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)) + })) +} diff --git a/proxy/masque/server_test.go b/proxy/masque/server_test.go new file mode 100644 index 000000000..9b526f764 --- /dev/null +++ b/proxy/masque/server_test.go @@ -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) + } +} diff --git a/proxy/wireguard/server.go b/proxy/wireguard/server.go index a9a0a73ca..f8bdf218b 100644 --- a/proxy/wireguard/server.go +++ b/proxy/wireguard/server.go @@ -320,7 +320,7 @@ func (s *Server) Start() error { return err } s.dev = dev - createForwarder(s.stack, s.HandleConnection) + CreateForwarder(s.stack, s.HandleConnection) return nil } diff --git a/proxy/wireguard/tun.go b/proxy/wireguard/tun.go index 68bad3ac6..b4bcdc371 100644 --- a/proxy/wireguard/tun.go +++ b/proxy/wireguard/tun.go @@ -49,7 +49,7 @@ func CalculateInterfaceName(name string) (tunName string) { return } -func createForwarder(gstack *stack.Stack, handler func(conn net.Conn, dest net.Destination)) { +func CreateForwarder(gstack *stack.Stack, handler func(conn net.Conn, dest net.Destination)) { gstack.SetPromiscuousMode(1, true) gstack.SetSpoofing(1, true) diff --git a/testing/scenarios/masque_test.go b/testing/scenarios/masque_test.go index 6f0f92d3d..95526a222 100644 --- a/testing/scenarios/masque_test.go +++ b/testing/scenarios/masque_test.go @@ -4,8 +4,10 @@ import ( "bufio" "bytes" "context" + "crypto/rand" gotls "crypto/tls" "crypto/x509" + "encoding/binary" go_errors "errors" "io" "net/http" @@ -30,6 +32,7 @@ import ( "github.com/xtls/xray-core/app/log" "github.com/xtls/xray-core/app/proxyman" + "github.com/xtls/xray-core/app/router" "github.com/xtls/xray-core/common" clog "github.com/xtls/xray-core/common/log" "github.com/xtls/xray-core/common/net" @@ -38,6 +41,7 @@ import ( "github.com/xtls/xray-core/common/serial" core "github.com/xtls/xray-core/core" "github.com/xtls/xray-core/proxy/dokodemo" + "github.com/xtls/xray-core/proxy/freedom" "github.com/xtls/xray-core/proxy/masque" "github.com/xtls/xray-core/proxy/wireguard" "github.com/xtls/xray-core/testing/servers/tcp" @@ -57,7 +61,7 @@ var ( const ( masqueEchoPort = 7 - masqueAuthorization = "Basic dTpw" + masqueAuthorization = "Basic dUBleGFtcGxlLmNvbTpw" ) func startMasqueServer(t *testing.T, h2 bool) (net.Port, [32]byte) { @@ -384,33 +388,65 @@ func TestMasqueHTTP2(t *testing.T) { testMasque(t, true) } -func testMasque(t *testing.T, h2 bool) { - serverPort, certHash := startMasqueServer(t, h2) - - tcpPort := tcp.PickPort() - tcp6Port := tcp.PickPort() - udpPort := udp.PickPort() - dokodemoTo := func(port net.Port, addr netip.Addr, network net.Network) *core.InboundHandlerConfig { - return &core.InboundHandlerConfig{ - ReceiverSettings: serial.ToTypedMessage(&proxyman.ReceiverConfig{ - PortList: &net.PortList{Range: []*net.PortRange{net.SinglePortRange(port)}}, - Listen: net.NewIPOrDomain(net.LocalHostIP), - }), - ProxySettings: serial.ToTypedMessage(&dokodemo.Config{ - RewriteAddress: net.NewIPOrDomain(net.IPAddress(addr.AsSlice())), - RewritePort: masqueEchoPort, - AllowedNetworks: []net.Network{network}, - }), - } +func masqueDokodemo(port net.Port, addr netip.Addr, network net.Network) *core.InboundHandlerConfig { + return &core.InboundHandlerConfig{ + ReceiverSettings: serial.ToTypedMessage(&proxyman.ReceiverConfig{ + PortList: &net.PortList{Range: []*net.PortRange{net.SinglePortRange(port)}}, + Listen: net.NewIPOrDomain(net.LocalHostIP), + }), + ProxySettings: serial.ToTypedMessage(&dokodemo.Config{ + RewriteAddress: net.NewIPOrDomain(net.IPAddress(addr.AsSlice())), + RewritePort: masqueEchoPort, + AllowedNetworks: []net.Network{network}, + }), } - tlsConfig := &tls.Config{ +} + +func masqueStreamSettings(tlsConfig *tls.Config, config *transmasque.Config) *internet.StreamConfig { + return &internet.StreamConfig{ + ProtocolName: "masque", + TransportSettings: []*internet.TransportConfig{ + { + ProtocolName: "masque", + Settings: serial.ToTypedMessage(config), + }, + }, + SecurityType: serial.GetMessageType(&tls.Config{}), + SecuritySettings: []*serial.TypedMessage{serial.ToTypedMessage(tlsConfig)}, + } +} + +func masqueClientTLS(certHash [32]byte, alpn ...string) *tls.Config { + return &tls.Config{ ServerName: "localhost", PinnedPeerCertSha256: [][]byte{certHash[:]}, + NextProtocol: alpn, } +} + +func masqueOutbound(serverPort net.Port, certHash [32]byte, h2 bool, authorization string) *core.OutboundHandlerConfig { + tlsConfig := masqueClientTLS(certHash) if h2 { tlsConfig.NextProtocol = []string{http2.NextProtoTLS} } - clientConfig := &core.Config{ + return &core.OutboundHandlerConfig{ + ProxySettings: serial.ToTypedMessage(&masque.ClientConfig{ + Server: &protocol.ServerEndpoint{ + Address: net.NewIPOrDomain(net.LocalHostIP), + Port: uint32(serverPort), + }, + }), + SenderSettings: serial.ToTypedMessage(&proxyman.SenderConfig{ + StreamSettings: masqueStreamSettings(tlsConfig, &transmasque.Config{ + Path: transmasque.DefaultPath, + Headers: map[string]string{"Authorization": authorization}, + }), + }), + } +} + +func masqueClientConfig(serverPort net.Port, certHash [32]byte, h2 bool, authorization string, tcpPort, tcp6Port, udpPort net.Port, v4, v6 netip.Addr) *core.Config { + return &core.Config{ App: []*serial.TypedMessage{ serial.ToTypedMessage(&log.Config{ ErrorLogLevel: clog.Severity_Debug, @@ -418,44 +454,17 @@ func testMasque(t *testing.T, h2 bool) { }), }, Inbound: []*core.InboundHandlerConfig{ - dokodemoTo(tcpPort, masqueServerV4, net.Network_TCP), - dokodemoTo(tcp6Port, masqueServerV6, net.Network_TCP), - dokodemoTo(udpPort, masqueServerV4, net.Network_UDP), + masqueDokodemo(tcpPort, v4, net.Network_TCP), + masqueDokodemo(tcp6Port, v6, net.Network_TCP), + masqueDokodemo(udpPort, v4, net.Network_UDP), }, Outbound: []*core.OutboundHandlerConfig{ - { - ProxySettings: serial.ToTypedMessage(&masque.ClientConfig{ - Server: &protocol.ServerEndpoint{ - Address: net.NewIPOrDomain(net.LocalHostIP), - Port: uint32(serverPort), - }, - }), - SenderSettings: serial.ToTypedMessage(&proxyman.SenderConfig{ - StreamSettings: &internet.StreamConfig{ - ProtocolName: "masque", - TransportSettings: []*internet.TransportConfig{ - { - ProtocolName: "masque", - Settings: serial.ToTypedMessage(&transmasque.Config{ - Path: transmasque.DefaultPath, - Headers: map[string]string{"Authorization": masqueAuthorization}, - }), - }, - }, - SecurityType: serial.GetMessageType(&tls.Config{}), - SecuritySettings: []*serial.TypedMessage{ - serial.ToTypedMessage(tlsConfig), - }, - }, - }), - }, + masqueOutbound(serverPort, certHash, h2, authorization), }, } +} - servers, err := InitializeServerConfigs(clientConfig) - common.Must(err) - defer CloseAllServers(servers) - +func testMasqueTraffic(t *testing.T, tcpPort, tcp6Port, udpPort net.Port) { var errg errgroup.Group for range 3 { errg.Go(testTCPConn(tcpPort, 1024*1024, time.Second*20)) @@ -466,3 +475,215 @@ func testMasque(t *testing.T, h2 bool) { t.Error(err) } } + +func testMasque(t *testing.T, h2 bool) { + serverPort, certHash := startMasqueServer(t, h2) + + tcpPort := tcp.PickPort() + tcp6Port := tcp.PickPort() + udpPort := udp.PickPort() + clientConfig := masqueClientConfig(serverPort, certHash, h2, masqueAuthorization, tcpPort, tcp6Port, udpPort, masqueServerV4, masqueServerV6) + + servers, err := InitializeServerConfigs(clientConfig) + common.Must(err) + defer CloseAllServers(servers) + + testMasqueTraffic(t, tcpPort, tcp6Port, udpPort) +} + +func masqueServerInbound(serverPort net.Port, certificate *tls.Certificate, alpn ...string) *core.InboundHandlerConfig { + return &core.InboundHandlerConfig{ + ReceiverSettings: serial.ToTypedMessage(&proxyman.ReceiverConfig{ + PortList: &net.PortList{Range: []*net.PortRange{net.SinglePortRange(serverPort)}}, + Listen: net.NewIPOrDomain(net.LocalHostIP), + StreamSettings: masqueStreamSettings(&tls.Config{ + Certificate: []*tls.Certificate{certificate}, + NextProtocol: alpn, + }, &transmasque.Config{Path: transmasque.DefaultPath}), + }), + ProxySettings: serial.ToTypedMessage(&masque.ServerConfig{ + Users: []*protocol.User{{ + Email: "u@example.com", + Account: serial.ToTypedMessage(&masque.Account{Password: "p"}), + }}, + Address: []string{"10.14.0.1/24", "fd14::1/64"}, + }), + } +} + +func masqueServerConfig(serverPort net.Port, certificate *tls.Certificate, h2 bool, tcpDest, udpDest net.Destination) *core.Config { + var alpn []string + if h2 { + alpn = []string{http2.NextProtoTLS} + } + redirect := func(tag string, dest net.Destination) *core.OutboundHandlerConfig { + return &core.OutboundHandlerConfig{ + Tag: tag, + ProxySettings: serial.ToTypedMessage(&freedom.Config{ + DestinationOverride: &freedom.DestinationOverride{ + Server: &protocol.ServerEndpoint{ + Address: net.NewIPOrDomain(dest.Address), + Port: uint32(dest.Port), + }, + }, + FinalRules: []*freedom.FinalRuleConfig{{Action: freedom.RuleAction_Allow}}, + }), + } + } + return &core.Config{ + App: []*serial.TypedMessage{ + serial.ToTypedMessage(&log.Config{ + ErrorLogLevel: clog.Severity_Debug, + ErrorLogType: log.LogType_Console, + }), + serial.ToTypedMessage(&router.Config{ + Rule: []*router.RoutingRule{ + {Networks: []net.Network{net.Network_TCP}, TargetTag: &router.RoutingRule_Tag{Tag: "tcp"}}, + {Networks: []net.Network{net.Network_UDP}, TargetTag: &router.RoutingRule_Tag{Tag: "udp"}}, + }, + }), + }, + Inbound: []*core.InboundHandlerConfig{ + masqueServerInbound(serverPort, certificate, alpn...), + }, + Outbound: []*core.OutboundHandlerConfig{ + redirect("tcp", tcpDest), + redirect("udp", udpDest), + }, + } +} + +func testMasqueServer(t *testing.T, h2 bool, authorization string) error { + tcpServer := tcp.Server{MsgProcessor: xor} + tcpDest, err := tcpServer.Start() + common.Must(err) + defer tcpServer.Close() + udpServer := udp.Server{MsgProcessor: xor} + udpDest, err := udpServer.Start() + common.Must(err) + defer udpServer.Close() + + ct, ctHash := cert.MustGenerate(nil, cert.CommonName("localhost")) + serverPort := udp.PickPort() + if h2 { + serverPort = tcp.PickPort() + } + tcpPort := tcp.PickPort() + tcp6Port := tcp.PickPort() + udpPort := udp.PickPort() + servers, err := InitializeServerConfigs( + masqueServerConfig(serverPort, tls.ParseCertificate(ct), h2, tcpDest, udpDest), + masqueClientConfig(serverPort, ctHash, h2, authorization, tcpPort, tcp6Port, udpPort, netip.MustParseAddr("192.0.2.1"), netip.MustParseAddr("2001:db8::1")), + ) + common.Must(err) + defer CloseAllServers(servers) + + if authorization != masqueAuthorization { + return testTCPConn(tcpPort, 1024, time.Second*5)() + } + testMasqueTraffic(t, tcpPort, tcp6Port, udpPort) + return nil +} + +func TestMasqueServer(t *testing.T) { + testMasqueServer(t, false, masqueAuthorization) +} + +func TestMasqueServerHTTP2(t *testing.T) { + testMasqueServer(t, true, masqueAuthorization) +} + +func TestMasqueServerRejectsWrongPassword(t *testing.T) { + for _, h2 := range []bool{false, true} { + if err := testMasqueServer(t, h2, "Basic dUBleGFtcGxlLmNvbTp3cm9uZw=="); err == nil { + t.Errorf("a wrong password got through (h2: %v)", h2) + } + } +} + +func masqueIPPacket(src, dst netip.Addr, payload []byte) []byte { + if src.Is4() { + p := make([]byte, 20, 20+len(payload)) + p[0] = 0x45 + binary.BigEndian.PutUint16(p[2:], uint16(20+len(payload))) + p[8] = 64 + p[9] = 253 + copy(p[12:], src.AsSlice()) + copy(p[16:], dst.AsSlice()) + return append(p, payload...) + } + p := make([]byte, 40, 40+len(payload)) + p[0] = 0x60 + binary.BigEndian.PutUint16(p[4:], uint16(len(payload))) + p[6] = 253 + p[7] = 64 + copy(p[8:], src.AsSlice()) + copy(p[24:], dst.AsSlice()) + return append(p, payload...) +} + +func masqueIPAddrs(p []byte) (src, dst netip.Addr) { + if p[0]>>4 == 4 { + return netip.AddrFrom4([4]byte(p[12:16])), netip.AddrFrom4([4]byte(p[16:20])) + } + return netip.AddrFrom16([16]byte(p[8:24])), netip.AddrFrom16([16]byte(p[24:40])) +} + +func TestMasqueServerClientToClient(t *testing.T) { + ct, ctHash := cert.MustGenerate(nil, cert.CommonName("localhost")) + serverPort := udp.PickPort() + servers, err := InitializeServerConfigs(&core.Config{ + Inbound: []*core.InboundHandlerConfig{ + masqueServerInbound(serverPort, tls.ParseCertificate(ct), http3.NextProtoH3, http2.NextProtoTLS), + }, + Outbound: []*core.OutboundHandlerConfig{ + {ProxySettings: serial.ToTypedMessage(&freedom.Config{})}, + }, + }) + common.Must(err) + defer CloseAllServers(servers) + + dial := func(alpn ...string) *transmasque.Conn { + streamSettings, err := internet.ToMemoryStreamConfig(masqueStreamSettings(masqueClientTLS(ctHash, alpn...), &transmasque.Config{ + Path: transmasque.DefaultPath, + Headers: map[string]string{"Authorization": masqueAuthorization}, + })) + common.Must(err) + ctx, cancel := context.WithTimeout(context.Background(), 10*time.Second) + defer cancel() + conn, err := transmasque.Dial(ctx, net.TCPDestination(net.LocalHostIP, serverPort), streamSettings) + if err != nil { + t.Fatal(err) + } + t.Cleanup(func() { conn.Close() }) + return conn.(*transmasque.Conn) + } + h3 := dial() + h2 := dial(http2.NextProtoTLS) + + for _, c := range []struct{ from, to *transmasque.Conn }{{h3, h2}, {h2, h3}} { + for i := range c.from.LocalAddrs() { + src, dst := c.from.LocalAddrs()[i], c.to.LocalAddrs()[i] + payload := make([]byte, 1000) + rand.Read(payload) + if _, err := c.from.Write(masqueIPPacket(src, dst, payload)); err != nil { + t.Fatal(err) + } + received := make(chan []byte, 1) + go func() { + b := make([]byte, 2048) + n, _ := c.to.Read(b) + received <- b[:n] + }() + select { + case p := <-received: + gotSrc, gotDst := masqueIPAddrs(p) + if gotSrc != src || gotDst != dst || !bytes.HasSuffix(p, payload) { + t.Fatalf("unexpected packet from %s to %s: %x", gotSrc, gotDst, p) + } + case <-time.After(5 * time.Second): + t.Fatalf("no packet from %s to %s", src, dst) + } + } + } +} diff --git a/transport/internet/masque/connectip/http2.go b/transport/internet/masque/connectip/http2.go index a83f513e5..be4b52d72 100644 --- a/transport/internet/masque/connectip/http2.go +++ b/transport/internet/masque/connectip/http2.go @@ -15,7 +15,7 @@ import ( "github.com/apernet/quic-go" ) -const maxBufferedRequestBody = 32 << 10 +const maxStreamBuffer = 32 << 10 type HTTP2ClientConn struct { roundTripper http.RoundTripper @@ -37,7 +37,7 @@ func (c *HTTP2ClientConn) Dial(req *Request) (*Conn, *http.Response, error) { ctx := httpReq.Context() streamCtx, cancel := context.WithCancel(context.WithoutCancel(ctx)) stop := context.AfterFunc(ctx, cancel) - body := newRequestBody() + body := NewStreamBuffer() r := httpReq.Clone(streamCtx) r.Header[":protocol"] = []string{requestProtocol} r.Body = body @@ -67,7 +67,7 @@ func (c *HTTP2ClientConn) Dial(req *Request) (*Conn, *http.Response, error) { type http2Stream struct { reader *bufio.Reader - body *requestBody + body *StreamBuffer rsp io.Closer cancel context.CancelFunc } @@ -86,7 +86,7 @@ func (s *http2Stream) abort() { s.rsp.Close() } -type requestBody struct { +type StreamBuffer struct { mu sync.Mutex cond sync.Cond buf []byte @@ -95,13 +95,13 @@ type requestBody struct { deadline time.Time } -func newRequestBody() *requestBody { - b := &requestBody{} +func NewStreamBuffer() *StreamBuffer { + b := &StreamBuffer{} b.cond.L = &b.mu return b } -func (b *requestBody) Read(p []byte) (int, error) { +func (b *StreamBuffer) Read(p []byte) (int, error) { b.mu.Lock() defer b.mu.Unlock() for len(b.buf) == 0 && !b.closed && b.err == nil { @@ -119,7 +119,7 @@ func (b *requestBody) Read(p []byte) (int, error) { return n, nil } -func (b *requestBody) Write(p []byte) (int, error) { +func (b *StreamBuffer) Write(p []byte) (int, error) { b.mu.Lock() defer b.mu.Unlock() for { @@ -130,7 +130,7 @@ func (b *requestBody) Write(p []byte) (int, error) { return 0, io.ErrClosedPipe case !b.deadline.IsZero() && !time.Now().Before(b.deadline): return 0, os.ErrDeadlineExceeded - case len(b.buf) < maxBufferedRequestBody: + case len(b.buf) < maxStreamBuffer: b.buf = append(b.buf, p...) b.cond.Broadcast() return len(p), nil @@ -139,7 +139,7 @@ func (b *requestBody) Write(p []byte) (int, error) { } } -func (b *requestBody) Close() error { +func (b *StreamBuffer) Close() error { b.mu.Lock() b.closed = true b.cond.Broadcast() @@ -147,7 +147,7 @@ func (b *requestBody) Close() error { return nil } -func (b *requestBody) CloseWithError(err error) { +func (b *StreamBuffer) CloseWithError(err error) { b.mu.Lock() if b.err == nil { b.err = err @@ -157,7 +157,7 @@ func (b *requestBody) CloseWithError(err error) { b.mu.Unlock() } -func (b *requestBody) SetWriteDeadline(t time.Time) error { +func (b *StreamBuffer) SetWriteDeadline(t time.Time) error { b.mu.Lock() b.deadline = t b.cond.Broadcast() diff --git a/transport/internet/masque/connectip/http2_test.go b/transport/internet/masque/connectip/http2_test.go index e5c9d3cd0..4ddc9c52f 100644 --- a/transport/internet/masque/connectip/http2_test.go +++ b/transport/internet/masque/connectip/http2_test.go @@ -109,7 +109,7 @@ func setupHTTP2Conns(t *testing.T) (client, server *Conn) { func newTestHTTP2Stream() (*http2Stream, *io.PipeWriter) { pr, pw := io.Pipe() - return &http2Stream{reader: bufio.NewReader(pr), body: newRequestBody(), rsp: pr, cancel: func() {}}, pw + return &http2Stream{reader: bufio.NewReader(pr), body: NewStreamBuffer(), rsp: pr, cancel: func() {}}, pw } func TestHTTP2Request(t *testing.T) { @@ -343,7 +343,7 @@ func TestHTTP2CloseUnblocksWrites(t *testing.T) { require.Eventually(t, func() bool { str.body.mu.Lock() defer str.body.mu.Unlock() - return len(str.body.buf) >= maxBufferedRequestBody + return len(str.body.buf) >= maxStreamBuffer }, 5*time.Second, time.Millisecond) closed := make(chan error, 1) @@ -416,7 +416,7 @@ func TestHTTP2WritesDatagramCapsules(t *testing.T) { func TestRequestBody(t *testing.T) { t.Run("coalesces writes", func(t *testing.T) { - b := newRequestBody() + b := NewStreamBuffer() for _, s := range []string{"foo", "bar", "baz"} { _, err := b.Write([]byte(s)) require.NoError(t, err) @@ -428,8 +428,8 @@ func TestRequestBody(t *testing.T) { }) t.Run("blocks writes while full", func(t *testing.T) { - b := newRequestBody() - _, err := b.Write(make([]byte, maxBufferedRequestBody)) + b := NewStreamBuffer() + _, err := b.Write(make([]byte, maxStreamBuffer)) require.NoError(t, err) written := make(chan struct{}) go func() { @@ -441,7 +441,7 @@ func TestRequestBody(t *testing.T) { t.Fatal("write did not block") case <-time.After(50 * time.Millisecond): } - _, err = b.Read(make([]byte, maxBufferedRequestBody)) + _, err = b.Read(make([]byte, maxStreamBuffer)) require.NoError(t, err) select { case <-written: @@ -451,7 +451,7 @@ func TestRequestBody(t *testing.T) { }) t.Run("close", func(t *testing.T) { - b := newRequestBody() + b := NewStreamBuffer() _, err := b.Write([]byte("foo")) require.NoError(t, err) require.NoError(t, b.Close()) @@ -463,8 +463,8 @@ func TestRequestBody(t *testing.T) { }) t.Run("write deadline", func(t *testing.T) { - b := newRequestBody() - _, err := b.Write(make([]byte, maxBufferedRequestBody)) + b := NewStreamBuffer() + _, err := b.Write(make([]byte, maxStreamBuffer)) require.NoError(t, err) writeErr := make(chan error, 1) go func() { @@ -481,11 +481,11 @@ func TestRequestBody(t *testing.T) { require.NoError(t, b.Close()) data, err := io.ReadAll(b) require.NoError(t, err) - require.Len(t, data, maxBufferedRequestBody) + require.Len(t, data, maxStreamBuffer) }) t.Run("close with error", func(t *testing.T) { - b := newRequestBody() + b := NewStreamBuffer() _, err := b.Write([]byte("foo")) require.NoError(t, err) b.CloseWithError(net.ErrClosed) diff --git a/transport/internet/masque/http2_server.go b/transport/internet/masque/http2_server.go new file mode 100644 index 000000000..98aaa25bf --- /dev/null +++ b/transport/internet/masque/http2_server.go @@ -0,0 +1,680 @@ +package masque + +import ( + "bufio" + "bytes" + "context" + go_errors "errors" + "io" + "maps" + "math" + "net" + "net/http" + "net/url" + "slices" + "strconv" + "strings" + "sync" + "sync/atomic" + "time" + + "github.com/xtls/xray-core/transport/internet/masque/connectip" + "golang.org/x/net/http2" + "golang.org/x/net/http2/hpack" +) + +const ( + http2MaxConcurrentStreams = 100 + http2HandshakeTimeout = 10 * time.Second + http2DefaultHeaderTable = 4096 +) + +var ( + errHTTP2BadPreface = go_errors.New("http2: invalid connection preface") + errHTTP2BadRequest = go_errors.New("http2: malformed request") + errHTTP2StreamClosed = go_errors.New("http2: stream closed") + errHTTP2RequestBodyClosed = go_errors.New("http2: request body closed") +) + +type connAddrsKey struct{} + +type connAddrs struct { + local net.Addr + remote net.Addr +} + +type http2ServerConn struct { + conn net.Conn + handler http.Handler + ctx context.Context + cancel context.CancelFunc + + wmu sync.Mutex + bw *bufio.Writer + fr *http2.Framer + hbuf bytes.Buffer + henc *hpack.Encoder + + lastFrame atomic.Int64 + + mu sync.Mutex + cond sync.Cond + err error + maxFrameSize uint32 + initialWindow int64 + connSendWindow int64 + connRecvWindow int64 + connRecvUnacked int64 + streams map[uint32]*http2ServerStream + lastStreamID uint32 +} + +type http2ServerStream struct { + c *http2ServerConn + id uint32 + ctx context.Context + cancel context.CancelFunc + header http.Header + out *connectip.StreamBuffer + sent chan struct{} + + sendWindow int64 + recvWindow int64 + recvUnacked int64 + recv bytes.Buffer + recvEnd bool + wroteHeader bool + sentEnd bool + resetErr error +} + +func serveHTTP2(ctx context.Context, conn net.Conn, handler http.Handler) { + ctx, cancel := context.WithCancel(ctx) + c := &http2ServerConn{ + conn: conn, + handler: handler, + ctx: context.WithValue(ctx, connAddrsKey{}, connAddrs{local: conn.LocalAddr(), remote: conn.RemoteAddr()}), + cancel: cancel, + bw: bufio.NewWriter(conn), + maxFrameSize: http2DefaultFrameSize, + initialWindow: http2DefaultWindow, + connSendWindow: http2DefaultWindow, + connRecvWindow: http2ConnectionWindow, + streams: make(map[uint32]*http2ServerStream), + } + c.cond.L = &c.mu + c.henc = hpack.NewEncoder(&c.hbuf) + c.henc.SetMaxDynamicTableSizeLimit(0) + c.lastFrame.Store(time.Now().UnixNano()) + + br := bufio.NewReader(conn) + preface := make([]byte, len(http2.ClientPreface)) + conn.SetReadDeadline(time.Now().Add(http2HandshakeTimeout)) + if _, err := io.ReadFull(br, preface); err != nil || string(preface) != http2.ClientPreface { + c.fail(errHTTP2BadPreface) + return + } + conn.SetReadDeadline(time.Time{}) + + c.fr = http2.NewFramer(c.bw, br) + c.fr.SetMaxReadFrameSize(http2DefaultFrameSize) + c.fr.ReadMetaHeaders = hpack.NewDecoder(http2DefaultHeaderTable, nil) + c.fr.MaxHeaderListSize = http2MaxHeaderListSize + if err := c.write(func(fr *http2.Framer) error { + if err := fr.WriteSettings( + http2.Setting{ID: http2.SettingMaxConcurrentStreams, Val: http2MaxConcurrentStreams}, + http2.Setting{ID: http2.SettingInitialWindowSize, Val: http2StreamWindow}, + http2.Setting{ID: http2.SettingMaxHeaderListSize, Val: http2MaxHeaderListSize}, + http2.Setting{ID: http2.SettingEnableConnectProtocol, Val: 1}, + ); err != nil { + return err + } + return fr.WriteWindowUpdate(0, http2ConnectionWindow-http2DefaultWindow) + }); err != nil { + c.fail(err) + return + } + go c.keepAlive() + c.readLoop() +} + +func (c *http2ServerConn) write(f func(*http2.Framer) error) error { + c.wmu.Lock() + defer c.wmu.Unlock() + if err := f(c.fr); err != nil { + return err + } + return c.bw.Flush() +} + +func (c *http2ServerConn) fail(err error) { + c.mu.Lock() + if c.err == nil { + c.err = err + } + streams := slices.Collect(maps.Values(c.streams)) + c.cond.Broadcast() + c.mu.Unlock() + for _, st := range streams { + st.out.CloseWithError(err) + st.cancel() + } + c.cancel() + c.conn.Close() +} + +func (c *http2ServerConn) keepAlive() { + ticker := time.NewTicker(http2KeepAlivePeriod) + defer ticker.Stop() + for { + select { + case <-c.ctx.Done(): + return + case <-ticker.C: + } + idle := time.Since(time.Unix(0, c.lastFrame.Load())) + if idle >= http2IdleTimeout { + c.fail(errHTTP2IdleTimeout) + return + } + if idle >= http2KeepAlivePeriod { + go c.write(func(fr *http2.Framer) error { + return fr.WritePing(false, [8]byte{}) + }) + } + } +} + +func (c *http2ServerConn) readLoop() { + for { + f, err := c.fr.ReadFrame() + if err != nil { + var streamErr http2.StreamError + if go_errors.As(err, &streamErr) { + if err := c.resetStream(streamErr); err != nil { + c.fail(err) + return + } + continue + } + c.fail(err) + return + } + c.lastFrame.Store(time.Now().UnixNano()) + if err := c.handleFrame(f); err != nil { + c.fail(err) + return + } + } +} + +func (c *http2ServerConn) handleFrame(f http2.Frame) error { + switch f := f.(type) { + case *http2.SettingsFrame: + if f.IsAck() { + return nil + } + if err := c.applySettings(f); err != nil { + return err + } + return c.write((*http2.Framer).WriteSettingsAck) + case *http2.PingFrame: + if f.IsAck() { + return nil + } + return c.write(func(fr *http2.Framer) error { + return fr.WritePing(true, f.Data) + }) + case *http2.WindowUpdateFrame: + return c.handleWindowUpdate(f) + case *http2.MetaHeadersFrame: + return c.handleHeaders(f) + case *http2.DataFrame: + return c.handleData(f) + case *http2.RSTStreamFrame: + c.abortStream(f.StreamID, http2.StreamError{StreamID: f.StreamID, Code: f.ErrCode}, false) + case *http2.PushPromiseFrame: + return http2.ConnectionError(http2.ErrCodeProtocol) + } + return nil +} + +func (c *http2ServerConn) applySettings(f *http2.SettingsFrame) error { + c.mu.Lock() + defer c.mu.Unlock() + defer c.cond.Broadcast() + return f.ForeachSetting(func(s http2.Setting) error { + if err := s.Valid(); err != nil { + return err + } + switch s.ID { + case http2.SettingMaxFrameSize: + c.maxFrameSize = s.Val + case http2.SettingInitialWindowSize: + delta := int64(s.Val) - c.initialWindow + c.initialWindow = int64(s.Val) + for _, st := range c.streams { + st.sendWindow += delta + if st.sendWindow > math.MaxInt32 { + return http2.ConnectionError(http2.ErrCodeFlowControl) + } + } + } + return nil + }) +} + +func (c *http2ServerConn) handleWindowUpdate(f *http2.WindowUpdateFrame) error { + c.mu.Lock() + defer c.mu.Unlock() + if f.StreamID == 0 { + c.connSendWindow += int64(f.Increment) + if c.connSendWindow > math.MaxInt32 { + return http2.ConnectionError(http2.ErrCodeFlowControl) + } + } else if st := c.streams[f.StreamID]; st != nil { + st.sendWindow += int64(f.Increment) + if st.sendWindow > math.MaxInt32 { + return http2.ConnectionError(http2.ErrCodeFlowControl) + } + } + c.cond.Broadcast() + return nil +} + +func (c *http2ServerConn) handleHeaders(f *http2.MetaHeadersFrame) error { + c.mu.Lock() + if st := c.streams[f.StreamID]; st != nil { + if !f.StreamEnded() { + c.mu.Unlock() + return http2.ConnectionError(http2.ErrCodeProtocol) + } + st.recvEnd = true + c.cond.Broadcast() + c.mu.Unlock() + return nil + } + if f.StreamID%2 == 0 || f.StreamID <= c.lastStreamID { + c.mu.Unlock() + return http2.ConnectionError(http2.ErrCodeProtocol) + } + c.lastStreamID = f.StreamID + refused := len(c.streams) >= http2MaxConcurrentStreams + c.mu.Unlock() + + if refused { + return c.write(func(fr *http2.Framer) error { + return fr.WriteRSTStream(f.StreamID, http2.ErrCodeRefusedStream) + }) + } + req, err := newHTTP2Request(f) + if err != nil { + return c.write(func(fr *http2.Framer) error { + return fr.WriteRSTStream(f.StreamID, http2.ErrCodeProtocol) + }) + } + + ctx, cancel := context.WithCancel(c.ctx) + st := &http2ServerStream{ + c: c, + id: f.StreamID, + ctx: ctx, + cancel: cancel, + header: make(http.Header), + out: connectip.NewStreamBuffer(), + sent: make(chan struct{}), + recvWindow: http2StreamWindow, + recvEnd: f.StreamEnded(), + } + req = req.WithContext(ctx) + req.RemoteAddr = c.conn.RemoteAddr().String() + if st.recvEnd { + req.Body = http.NoBody + req.ContentLength = 0 + } else { + req.Body = &http2RequestBody{st: st} + req.ContentLength = -1 + } + + c.mu.Lock() + if c.err != nil { + err := c.err + c.mu.Unlock() + cancel() + return err + } + st.sendWindow = c.initialWindow + c.streams[f.StreamID] = st + c.mu.Unlock() + + go c.serveStream(st, req) + return nil +} + +func newHTTP2Request(f *http2.MetaHeadersFrame) (*http.Request, error) { + method := f.PseudoValue("method") + scheme := f.PseudoValue("scheme") + authority := f.PseudoValue("authority") + path := f.PseudoValue("path") + protocol := f.PseudoValue("protocol") + if method == "" || (protocol != "" && method != http.MethodConnect) { + return nil, errHTTP2BadRequest + } + var u *url.URL + if method == http.MethodConnect && protocol == "" { + if authority == "" || path != "" || scheme != "" { + return nil, errHTTP2BadRequest + } + u = &url.URL{Host: authority} + } else { + if path == "" || scheme == "" { + return nil, errHTTP2BadRequest + } + var err error + if u, err = url.ParseRequestURI(path); err != nil { + return nil, errHTTP2BadRequest + } + } + header := make(http.Header) + for _, hf := range f.RegularFields() { + header.Add(hf.Name, hf.Value) + } + if protocol != "" { + header.Set(":protocol", protocol) + } + if authority == "" { + authority = header.Get("Host") + } + return &http.Request{ + Method: method, + URL: u, + Proto: "HTTP/2.0", + ProtoMajor: 2, + Header: header, + Host: authority, + RequestURI: path, + }, nil +} + +func (c *http2ServerConn) handleData(f *http2.DataFrame) error { + size := int64(f.Length) + c.mu.Lock() + c.connRecvWindow -= size + if c.connRecvWindow < 0 { + c.mu.Unlock() + return http2.ConnectionError(http2.ErrCodeFlowControl) + } + st := c.streams[f.StreamID] + if st == nil || st.resetErr != nil || st.recvEnd { + idle := st == nil && f.StreamID > c.lastStreamID + c.connRecvWindow += size + c.mu.Unlock() + if idle { + return http2.ConnectionError(http2.ErrCodeProtocol) + } + if size == 0 { + return nil + } + return c.write(func(fr *http2.Framer) error { + return fr.WriteWindowUpdate(0, uint32(size)) + }) + } + st.recvWindow -= size + if st.recvWindow < 0 { + c.mu.Unlock() + return http2.ConnectionError(http2.ErrCodeFlowControl) + } + st.recv.Write(f.Data()) + padding := size - int64(len(f.Data())) + st.recvUnacked += padding + c.connRecvUnacked += padding + if f.StreamEnded() { + st.recvEnd = true + } + c.cond.Broadcast() + c.mu.Unlock() + return nil +} + +func (c *http2ServerConn) resetStream(streamErr http2.StreamError) error { + c.mu.Lock() + if streamErr.StreamID%2 == 1 && streamErr.StreamID > c.lastStreamID { + c.lastStreamID = streamErr.StreamID + } + c.mu.Unlock() + c.abortStream(streamErr.StreamID, streamErr, false) + return c.write(func(fr *http2.Framer) error { + return fr.WriteRSTStream(streamErr.StreamID, streamErr.Code) + }) +} + +func (c *http2ServerConn) abortStream(id uint32, err error, reset bool) { + c.mu.Lock() + st := c.streams[id] + if st == nil || st.resetErr != nil { + c.mu.Unlock() + return + } + st.resetErr = err + c.cond.Broadcast() + c.mu.Unlock() + st.out.CloseWithError(err) + st.cancel() + if reset { + go c.write(func(fr *http2.Framer) error { + return fr.WriteRSTStream(id, http2.ErrCodeCancel) + }) + } +} + +func (c *http2ServerConn) serveStream(st *http2ServerStream, req *http.Request) { + go st.sendLoop() + c.handler.ServeHTTP(&http2ResponseWriter{st: st}, req) + st.writeHeader(http.StatusOK) + st.out.Close() + <-st.sent + + c.mu.Lock() + finish := st.resetErr == nil && !st.sentEnd && c.err == nil + refuse := finish && !st.recvEnd + st.sentEnd = true + delete(c.streams, st.id) + c.cond.Broadcast() + c.mu.Unlock() + st.cancel() + + if finish { + if err := c.write(func(fr *http2.Framer) error { + if err := fr.WriteData(st.id, true, nil); err != nil { + return err + } + if refuse { + return fr.WriteRSTStream(st.id, http2.ErrCodeNo) + } + return nil + }); err != nil { + c.fail(err) + } + } +} + +func (st *http2ServerStream) writeHeader(code int) { + c := st.c + c.mu.Lock() + if st.wroteHeader || st.resetErr != nil || c.err != nil { + c.mu.Unlock() + return + } + st.wroteHeader = true + header := st.header.Clone() + maxFrameSize := int(c.maxFrameSize) + c.mu.Unlock() + + if err := c.write(func(fr *http2.Framer) error { + c.hbuf.Reset() + c.henc.WriteField(hpack.HeaderField{Name: ":status", Value: strconv.Itoa(code)}) + for _, k := range slices.Sorted(maps.Keys(header)) { + name := strings.ToLower(k) + switch name { + case "connection", "proxy-connection", "keep-alive", "transfer-encoding", "upgrade": + continue + } + for _, v := range header[k] { + c.henc.WriteField(hpack.HeaderField{Name: name, Value: v}) + } + } + block := c.hbuf.Bytes() + for first := true; first || len(block) > 0; first = false { + chunk := block[:min(len(block), maxFrameSize)] + block = block[len(chunk):] + var err error + if first { + err = fr.WriteHeaders(http2.HeadersFrameParam{StreamID: st.id, BlockFragment: chunk, EndHeaders: len(block) == 0}) + } else { + err = fr.WriteContinuation(st.id, len(block) == 0, chunk) + } + if err != nil { + return err + } + } + return nil + }); err != nil { + c.fail(err) + } +} + +func (st *http2ServerStream) sendLoop() { + defer close(st.sent) + c := st.c + buf := make([]byte, http2DefaultFrameSize) + for { + n, err := st.out.Read(buf) + for data := buf[:n]; len(data) > 0; { + allowed, err := st.awaitSendWindow(len(data)) + if err != nil { + st.out.CloseWithError(err) + return + } + if err := c.write(func(fr *http2.Framer) error { + return fr.WriteData(st.id, false, data[:allowed]) + }); err != nil { + c.fail(err) + return + } + data = data[allowed:] + } + if err != nil { + return + } + } +} + +func (st *http2ServerStream) awaitSendWindow(n int) (int, error) { + c := st.c + c.mu.Lock() + defer c.mu.Unlock() + for { + switch { + case st.resetErr != nil: + return 0, st.resetErr + case c.err != nil: + return 0, c.err + case st.sentEnd: + return 0, errHTTP2StreamClosed + } + if window := min(c.connSendWindow, st.sendWindow); window > 0 { + n = int(min(int64(n), window, int64(c.maxFrameSize))) + c.connSendWindow -= int64(n) + st.sendWindow -= int64(n) + return n, nil + } + c.cond.Wait() + } +} + +type http2ResponseWriter struct { + st *http2ServerStream +} + +func (w *http2ResponseWriter) Header() http.Header { return w.st.header } + +func (w *http2ResponseWriter) WriteHeader(code int) { w.st.writeHeader(code) } + +func (w *http2ResponseWriter) Write(p []byte) (int, error) { + w.st.writeHeader(http.StatusOK) + return w.st.out.Write(p) +} + +func (w *http2ResponseWriter) Flush() { w.st.writeHeader(http.StatusOK) } + +func (w *http2ResponseWriter) SetWriteDeadline(t time.Time) error { + return w.st.out.SetWriteDeadline(t) +} + +type http2RequestBody struct { + st *http2ServerStream +} + +func (b *http2RequestBody) Read(p []byte) (int, error) { + st := b.st + c := st.c + c.mu.Lock() + for st.recv.Len() == 0 && !st.recvEnd && st.resetErr == nil && c.err == nil { + c.cond.Wait() + } + if st.recv.Len() == 0 { + err := io.EOF + switch { + case st.resetErr != nil: + err = st.resetErr + case c.err != nil && !st.recvEnd: + err = c.err + } + c.mu.Unlock() + return 0, err + } + n, _ := st.recv.Read(p) + st.recvUnacked += int64(n) + c.connRecvUnacked += int64(n) + var streamUpdate, connUpdate int64 + if st.recvUnacked >= http2WindowUpdateSize && !st.recvEnd { + streamUpdate = st.recvUnacked + st.recvUnacked = 0 + st.recvWindow += streamUpdate + } + if c.connRecvUnacked >= http2WindowUpdateSize { + connUpdate = c.connRecvUnacked + c.connRecvUnacked = 0 + c.connRecvWindow += connUpdate + } + c.mu.Unlock() + + if streamUpdate > 0 || connUpdate > 0 { + if err := c.write(func(fr *http2.Framer) error { + if connUpdate > 0 { + if err := fr.WriteWindowUpdate(0, uint32(connUpdate)); err != nil { + return err + } + } + if streamUpdate > 0 { + return fr.WriteWindowUpdate(st.id, uint32(streamUpdate)) + } + return nil + }); err != nil { + c.fail(err) + } + } + return n, nil +} + +func (b *http2RequestBody) Close() error { + st := b.st + c := st.c + c.mu.Lock() + done := st.recvEnd + c.mu.Unlock() + if !done { + c.abortStream(st.id, errHTTP2RequestBodyClosed, true) + } + return nil +} diff --git a/transport/internet/masque/http2_server_test.go b/transport/internet/masque/http2_server_test.go new file mode 100644 index 000000000..510e37438 --- /dev/null +++ b/transport/internet/masque/http2_server_test.go @@ -0,0 +1,281 @@ +package masque + +import ( + "bytes" + "context" + "crypto/rand" + "crypto/sha256" + "io" + "net" + "net/http" + "sync" + "testing" + "time" + + "github.com/stretchr/testify/assert" + "github.com/stretchr/testify/require" + "golang.org/x/net/http2" + "golang.org/x/net/http2/hpack" +) + +func serveHTTP2Pipe(t *testing.T, handler http.Handler) net.Conn { + t.Helper() + client, server := tcpPipe(t) + done := make(chan struct{}) + go func() { + serveHTTP2(context.Background(), server, handler) + close(done) + }() + t.Cleanup(func() { + client.Close() + server.Close() + <-done + }) + return client +} + +type http2ClientPeer struct { + t *testing.T + conn net.Conn + fr *http2.Framer + hbuf bytes.Buffer + henc *hpack.Encoder +} + +func newHTTP2ClientPeer(t *testing.T, handler http.Handler) (*http2ClientPeer, []http2.Setting) { + t.Helper() + conn := serveHTTP2Pipe(t, handler) + p := &http2ClientPeer{t: t, conn: conn, fr: http2.NewFramer(conn, conn)} + p.henc = hpack.NewEncoder(&p.hbuf) + p.fr.ReadMetaHeaders = hpack.NewDecoder(4096, nil) + _, err := io.WriteString(conn, http2.ClientPreface) + require.NoError(t, err) + require.NoError(t, p.fr.WriteSettings()) + + f := p.readFrame() + require.IsType(t, &http2.SettingsFrame{}, f) + var settings []http2.Setting + f.(*http2.SettingsFrame).ForeachSetting(func(s http2.Setting) error { + settings = append(settings, s) + return nil + }) + f = p.readFrame() + require.IsType(t, &http2.WindowUpdateFrame{}, f) + require.Equal(t, uint32(http2ConnectionWindow-http2DefaultWindow), f.(*http2.WindowUpdateFrame).Increment) + f = p.readFrame() + require.True(t, f.(*http2.SettingsFrame).IsAck()) + require.NoError(t, p.fr.WriteSettingsAck()) + return p, settings +} + +func (p *http2ClientPeer) readFrame() http2.Frame { + p.t.Helper() + p.conn.SetReadDeadline(time.Now().Add(5 * time.Second)) + f, err := p.fr.ReadFrame() + require.NoError(p.t, err) + return f +} + +func (p *http2ClientPeer) writeHeaders(streamID uint32, endStream bool, fields ...string) { + p.t.Helper() + p.hbuf.Reset() + for i := 0; i < len(fields); i += 2 { + require.NoError(p.t, p.henc.WriteField(hpack.HeaderField{Name: fields[i], Value: fields[i+1]})) + } + require.NoError(p.t, p.fr.WriteHeaders(http2.HeadersFrameParam{ + StreamID: streamID, + BlockFragment: p.hbuf.Bytes(), + EndHeaders: true, + EndStream: endStream, + })) +} + +func (p *http2ClientPeer) writeConnect(streamID uint32) { + p.t.Helper() + p.writeHeaders(streamID, false, + ":method", "CONNECT", + ":protocol", "connect-ip", + ":scheme", "https", + ":authority", "proxy.example", + ":path", "/.well-known/masque/ip/*/*/", + "capsule-protocol", "?1", + ) +} + +func TestHTTP2ServerSettings(t *testing.T) { + _, settings := newHTTP2ClientPeer(t, http.NotFoundHandler()) + require.Equal(t, []http2.Setting{ + {ID: http2.SettingMaxConcurrentStreams, Val: http2MaxConcurrentStreams}, + {ID: http2.SettingInitialWindowSize, Val: http2StreamWindow}, + {ID: http2.SettingMaxHeaderListSize, Val: http2MaxHeaderListSize}, + {ID: http2.SettingEnableConnectProtocol, Val: 1}, + }, settings) +} + +func TestHTTP2ServerRoundTrip(t *testing.T) { + handler := http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) { + assert.Equal(t, http.MethodConnect, r.Method) + assert.Equal(t, "connect-ip", r.Header.Get(":protocol")) + assert.Equal(t, "?1", r.Header.Get("Capsule-Protocol")) + assert.Equal(t, "Basic dTpw", r.Header.Get("Authorization")) + assert.Equal(t, 2, r.ProtoMajor) + assert.Equal(t, "proxy.example", r.Host) + assert.Equal(t, "/.well-known/masque/ip/*/*/", r.URL.Path) + w.Header().Set("Capsule-Protocol", "?1") + w.WriteHeader(http.StatusOK) + assert.NoError(t, http.NewResponseController(w).Flush()) + _, err := io.Copy(w, r.Body) + assert.NoError(t, err) + }) + cc, err := newHTTP2ClientConn(serveHTTP2Pipe(t, handler)) + require.NoError(t, err) + defer cc.Close() + + pr, pw := io.Pipe() + rsp, err := cc.RoundTrip(connectRequest(t, context.Background(), pr)) + require.NoError(t, err) + require.Equal(t, http.StatusOK, rsp.StatusCode) + require.Equal(t, "?1", rsp.Header.Get("Capsule-Protocol")) + + payload := make([]byte, 3*http2ConnectionWindow/2) + rand.Read(payload) + go func() { + pw.Write(payload) + pw.Close() + }() + echoed := sha256.New() + n, err := io.Copy(echoed, rsp.Body) + require.NoError(t, err) + require.Equal(t, int64(len(payload)), n) + require.Equal(t, sha256.Sum256(payload), [32]byte(echoed.Sum(nil))) +} + +func TestHTTP2ServerStatus(t *testing.T) { + p, _ := newHTTP2ClientPeer(t, http.NotFoundHandler()) + p.writeConnect(1) + f := p.readFrame() + require.IsType(t, &http2.MetaHeadersFrame{}, f) + require.Equal(t, "404", f.(*http2.MetaHeadersFrame).PseudoValue("status")) + var body []byte + for { + f = p.readFrame() + require.IsType(t, &http2.DataFrame{}, f) + body = append(body, f.(*http2.DataFrame).Data()...) + if f.(*http2.DataFrame).StreamEnded() { + break + } + } + require.Equal(t, "404 page not found\n", string(body)) + f = p.readFrame() + require.IsType(t, &http2.RSTStreamFrame{}, f) + require.Equal(t, http2.ErrCodeNo, f.(*http2.RSTStreamFrame).ErrCode) +} + +func TestHTTP2ServerMalformedRequests(t *testing.T) { + for _, tc := range []struct { + name string + fields []string + }{ + {"no method", []string{":scheme", "https", ":path", "/", ":authority", "proxy.example"}}, + {"no path", []string{":method", "GET", ":scheme", "https", ":authority", "proxy.example"}}, + {"protocol without CONNECT", []string{":method", "GET", ":protocol", "connect-ip", ":scheme", "https", ":path", "/", ":authority", "proxy.example"}}, + {"plain CONNECT with a path", []string{":method", "CONNECT", ":path", "/", ":authority", "proxy.example"}}, + } { + t.Run(tc.name, func(t *testing.T) { + p, _ := newHTTP2ClientPeer(t, http.HandlerFunc(func(http.ResponseWriter, *http.Request) { + t.Error("the handler saw a malformed request") + })) + p.writeHeaders(1, false, tc.fields...) + f := p.readFrame() + require.IsType(t, &http2.RSTStreamFrame{}, f) + require.Equal(t, http2.ErrCodeProtocol, f.(*http2.RSTStreamFrame).ErrCode) + }) + } +} + +func TestHTTP2ServerRefusesExtraStreams(t *testing.T) { + release := make(chan struct{}) + var started sync.WaitGroup + started.Add(http2MaxConcurrentStreams) + p, _ := newHTTP2ClientPeer(t, http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) { + started.Done() + <-release + })) + defer close(release) + for i := range http2MaxConcurrentStreams { + p.writeConnect(uint32(2*i + 1)) + } + started.Wait() + p.writeConnect(2*http2MaxConcurrentStreams + 1) + f := p.readFrame() + require.IsType(t, &http2.RSTStreamFrame{}, f) + require.Equal(t, uint32(2*http2MaxConcurrentStreams+1), f.Header().StreamID) + require.Equal(t, http2.ErrCodeRefusedStream, f.(*http2.RSTStreamFrame).ErrCode) +} + +func TestHTTP2ServerClientReset(t *testing.T) { + readErr := make(chan error, 1) + canceled := make(chan struct{}) + p, _ := newHTTP2ClientPeer(t, http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) { + w.WriteHeader(http.StatusOK) + _, err := r.Body.Read(make([]byte, 1)) + readErr <- err + <-r.Context().Done() + close(canceled) + })) + p.writeConnect(1) + f := p.readFrame() + require.Equal(t, "200", f.(*http2.MetaHeadersFrame).PseudoValue("status")) + require.NoError(t, p.fr.WriteRSTStream(1, http2.ErrCodeCancel)) + select { + case err := <-readErr: + require.Equal(t, http2.StreamError{StreamID: 1, Code: http2.ErrCodeCancel}, err) + case <-time.After(5 * time.Second): + t.Fatal("the body read did not fail after RST_STREAM") + } + select { + case <-canceled: + case <-time.After(5 * time.Second): + t.Fatal("the request context was not canceled") + } +} + +func TestHTTP2ServerAnswersPings(t *testing.T) { + p, _ := newHTTP2ClientPeer(t, http.NotFoundHandler()) + data := [8]byte{8, 7, 6, 5, 4, 3, 2, 1} + require.NoError(t, p.fr.WritePing(false, data)) + f := p.readFrame() + require.IsType(t, &http2.PingFrame{}, f) + require.True(t, f.(*http2.PingFrame).IsAck()) + require.Equal(t, data, f.(*http2.PingFrame).Data) +} + +func TestHTTP2ServerRejectsOverflow(t *testing.T) { + p, _ := newHTTP2ClientPeer(t, http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) { + w.WriteHeader(http.StatusOK) + <-r.Context().Done() + })) + p.writeConnect(1) + p.readFrame() + chunk := make([]byte, http2DefaultFrameSize) + go func() { + for range http2StreamWindow/len(chunk) + 1 { + if p.fr.WriteData(1, false, chunk) != nil { + return + } + } + }() + p.conn.SetReadDeadline(time.Now().Add(5 * time.Second)) + _, err := io.Copy(io.Discard, p.conn) + require.NoError(t, err) +} + +func TestHTTP2ServerBadPreface(t *testing.T) { + conn := serveHTTP2Pipe(t, http.NotFoundHandler()) + _, err := io.WriteString(conn, "GET / HTTP/1.1\r\nHost: example.com\r\n\r\n") + require.NoError(t, err) + conn.SetReadDeadline(time.Now().Add(5 * time.Second)) + n, err := io.Copy(io.Discard, conn) + require.NoError(t, err) + require.Zero(t, n) +} diff --git a/transport/internet/masque/hub.go b/transport/internet/masque/hub.go new file mode 100644 index 000000000..dec37db38 --- /dev/null +++ b/transport/internet/masque/hub.go @@ -0,0 +1,422 @@ +package masque + +import ( + "context" + "crypto/rand" + gotls "crypto/tls" + go_errors "errors" + "io" + "maps" + "net/http" + "net/url" + "runtime" + "slices" + "strings" + "sync" + "time" + + "github.com/apernet/quic-go" + "github.com/apernet/quic-go/http3" + "github.com/xtls/xray-core/common" + "github.com/xtls/xray-core/common/errors" + "github.com/xtls/xray-core/common/net" + "github.com/xtls/xray-core/transport/internet" + "github.com/xtls/xray-core/transport/internet/hysteria/congestion" + "github.com/xtls/xray-core/transport/internet/hysteria/congestion/bbr" + "github.com/xtls/xray-core/transport/internet/masque/connectip" + "github.com/xtls/xray-core/transport/internet/tls" + "golang.org/x/net/http2" +) + +type Listener struct { + path pathMatcher + addConn internet.ConnHandler + ctx context.Context + cancel context.CancelFunc + + quicServer *http3.Server + quicListener *quic.Listener + transport *quic.Transport + pktConn net.PacketConn + tcpListener net.Listener + + mu sync.Mutex + conns map[net.Conn]struct{} +} + +func serverVersions(config *tls.Config) (h2, h3 bool) { + h2 = slices.Contains(config.NextProtocol, http2.NextProtoTLS) + h3 = slices.Contains(config.NextProtocol, http3.NextProtoH3) || !h2 + return h2, h3 +} + +func Listen(ctx context.Context, address net.Address, port net.Port, streamSettings *internet.MemoryStreamConfig, handler internet.ConnHandler) (internet.Listener, error) { + if address.Family().IsDomain() { + return nil, errors.New("address is domain") + } + tlsConfig := tls.ConfigFromStreamSettings(streamSettings) + if tlsConfig == nil { + return nil, errors.New("tls config is nil") + } + config := streamSettings.ProtocolSettings.(*Config) + path, err := newPathMatcher(config.Path) + if err != nil { + return nil, err + } + + l := &Listener{ + path: path, + addConn: handler, + conns: make(map[net.Conn]struct{}), + } + l.ctx, l.cancel = context.WithCancel(context.Background()) + h2, h3 := serverVersions(tlsConfig) + if h3 { + if err := l.listenHTTP3(address, port, streamSettings, tlsConfig); err != nil { + l.Close() + return nil, err + } + errors.LogInfo(ctx, "listening UDP for MASQUE over HTTP/3 on ", address, ":", port) + } + if h2 { + if err := l.listenHTTP2(ctx, address, port, streamSettings, tlsConfig); err != nil { + l.Close() + return nil, err + } + errors.LogInfo(ctx, "listening TCP for MASQUE over HTTP/2 on ", address, ":", port) + } + return l, nil +} + +func (l *Listener) listenHTTP3(address net.Address, port net.Port, streamSettings *internet.MemoryStreamConfig, tlsConfig *tls.Config) error { + quicParams := streamSettings.QuicParams + if quicParams == nil { + quicParams = &internet.QuicParams{ + BbrProfile: string(bbr.ProfileStandard), + } + } + switch quicParams.Congestion { + case "", "reno", "bbr", "brutal", "force-brutal": + default: + return errors.New("unknown congestion control: ", quicParams.Congestion) + } + quicConfig := &quic.Config{ + InitialStreamReceiveWindow: quicParams.InitStreamReceiveWindow, + MaxStreamReceiveWindow: quicParams.MaxStreamReceiveWindow, + InitialConnectionReceiveWindow: quicParams.InitConnReceiveWindow, + MaxConnectionReceiveWindow: quicParams.MaxConnReceiveWindow, + MaxIdleTimeout: time.Duration(quicParams.MaxIdleTimeout) * time.Second, + KeepAlivePeriod: time.Duration(quicParams.KeepAlivePeriod) * time.Second, + MaxIncomingStreams: quicParams.MaxIncomingStreams, + InitialPacketSize: initialPacketSize, + DisablePathMTUDiscovery: quicParams.DisablePathMtuDiscovery || (runtime.GOOS != "linux" && runtime.GOOS != "windows" && runtime.GOOS != "darwin"), + EnableDatagrams: true, + DisablePathManager: true, + } + if quicParams.MaxIdleTimeout == 0 { + quicConfig.MaxIdleTimeout = 30 * time.Second + } + + udpAddr := &net.UDPAddr{IP: address.IP(), Port: int(port)} + var err error + if streamSettings.FinalMask != nil { + l.pktConn, err = streamSettings.FinalMask.ListenPacket(context.Background(), udpAddr) + } else { + l.pktConn, err = internet.ListenSystemPacket(context.Background(), udpAddr, streamSettings.SocketSettings) + } + if err != nil { + return errors.New("failed to listen UDP on ", address, ":", port).Base(err) + } + var resetKey *quic.StatelessResetKey + if !quicParams.DisableStatelessReset { + resetKey = &quic.StatelessResetKey{} + common.Must2(rand.Read(resetKey[:])) + } + l.transport = &quic.Transport{Conn: l.pktConn, DisableGSO: quicParams.DisableGSO, StatelessResetKey: resetKey} + + gotlsConfig := tlsConfig.GetTLSConfig() + gotlsConfig.NextProtos = []string{http3.NextProtoH3} + l.quicListener, err = l.transport.Listen(gotlsConfig, quicConfig) + if err != nil { + return err + } + l.quicServer = &http3.Server{ + Handler: l, + EnableDatagrams: true, + ConnContext: func(ctx context.Context, conn *quic.Conn) context.Context { + switch quicParams.Congestion { + case "reno": + case "", "bbr", "brutal": + congestion.UseBBR(conn, bbr.Profile(quicParams.BbrProfile)) + case "force-brutal": + congestion.UseBrutal(conn, quicParams.BrutalUp, quicParams.BrutalDisableLossCompensation) + } + return context.WithValue(ctx, connAddrsKey{}, connAddrs{local: conn.LocalAddr(), remote: conn.RemoteAddr()}) + }, + } + go func() { + if err := l.quicServer.ServeListener(l.quicListener); err != nil && !go_errors.Is(err, quic.ErrServerClosed) && !go_errors.Is(err, http.ErrServerClosed) { + errors.LogErrorInner(context.Background(), err, "failed to serve MASQUE over HTTP/3") + } + }() + return nil +} + +func (l *Listener) listenHTTP2(ctx context.Context, address net.Address, port net.Port, streamSettings *internet.MemoryStreamConfig, tlsConfig *tls.Config) error { + tcpAddr := &net.TCPAddr{IP: address.IP(), Port: int(port)} + var err error + if streamSettings.FinalMask != nil { + l.tcpListener, err = streamSettings.FinalMask.Listen(ctx, tcpAddr) + } else { + l.tcpListener, err = internet.ListenSystem(ctx, tcpAddr, streamSettings.SocketSettings) + } + if err != nil { + return errors.New("failed to listen TCP on ", address, ":", port).Base(err) + } + gotlsConfig := tlsConfig.GetTLSConfig() + gotlsConfig.NextProtos = []string{http2.NextProtoTLS} + go l.acceptHTTP2(gotlsConfig) + return nil +} + +func (l *Listener) acceptHTTP2(config *gotls.Config) { + for { + conn, err := l.tcpListener.Accept() + if err != nil { + if l.ctx.Err() != nil || strings.Contains(err.Error(), "closed") { + return + } + errors.LogWarningInner(context.Background(), err, "failed to accept MASQUE connections") + if strings.Contains(err.Error(), "too many") { + time.Sleep(500 * time.Millisecond) + } + continue + } + go l.serveHTTP2Conn(conn, config) + } +} + +func (l *Listener) serveHTTP2Conn(conn net.Conn, config *gotls.Config) { + tlsConn := tls.Server(conn, config).(*tls.Conn) + if !l.track(tlsConn, true) { + tlsConn.Close() + return + } + defer l.track(tlsConn, false) + ctx, cancel := context.WithTimeout(l.ctx, http2HandshakeTimeout) + err := tlsConn.HandshakeContext(ctx) + cancel() + if err != nil { + errors.LogDebugInner(context.Background(), err, "MASQUE: TLS handshake failed") + tlsConn.Close() + return + } + if protocol := tlsConn.NegotiatedProtocol(); protocol != http2.NextProtoTLS { + errors.LogDebug(context.Background(), "MASQUE: the client negotiated ", protocol, " instead of h2") + tlsConn.Close() + return + } + serveHTTP2(l.ctx, tlsConn, l) +} + +func (l *Listener) track(conn net.Conn, add bool) bool { + l.mu.Lock() + defer l.mu.Unlock() + if !add { + delete(l.conns, conn) + return true + } + if l.ctx.Err() != nil { + return false + } + l.conns[conn] = struct{}{} + return true +} + +func (l *Listener) Addr() net.Addr { + if l.tcpListener != nil { + return l.tcpListener.Addr() + } + return l.quicListener.Addr() +} + +func (l *Listener) Close() error { + l.cancel() + var errs []error + if l.quicServer != nil { + errs = append(errs, l.quicServer.Close()) + } + if l.quicListener != nil { + errs = append(errs, l.quicListener.Close()) + } + if l.transport != nil { + errs = append(errs, l.transport.Close()) + } + if l.pktConn != nil { + errs = append(errs, l.pktConn.Close()) + } + if l.tcpListener != nil { + errs = append(errs, l.tcpListener.Close()) + } + l.mu.Lock() + for conn := range l.conns { + conn.Close() + } + l.mu.Unlock() + return errors.Combine(errs...) +} + +func (l *Listener) ServeHTTP(w http.ResponseWriter, r *http.Request) { + if !l.path.match(r.URL) { + w.WriteHeader(http.StatusNotFound) + return + } + request, err := connectip.ParseProxyRequest(r) + if err != nil { + status := http.StatusBadRequest + if perr, ok := go_errors.AsType[*connectip.ProxyRequestParseError](err); ok { + status = perr.HTTPStatus + } + w.WriteHeader(status) + return + } + addrs, _ := r.Context().Value(connAddrsKey{}).(connAddrs) + conn := &ServerConn{ + w: w, + request: r, + proxyRequest: request, + local: addrs.local, + remote: addrs.remote, + done: make(chan struct{}), + } + l.addConn(conn) + select { + case <-conn.done: + case <-l.ctx.Done(): + conn.Close() + } +} + +type pathMatcher struct { + path string + query url.Values +} + +func newPathMatcher(path string) (pathMatcher, error) { + u, err := url.ParseRequestURI(path) + if err != nil || !strings.HasPrefix(u.Path, "/") { + return pathMatcher{}, errors.New("invalid path: ", path) + } + return pathMatcher{path: u.Path, query: u.Query()}, nil +} + +func (m pathMatcher) match(u *url.URL) bool { + return u.Path == m.path && maps.EqualFunc(u.Query(), m.query, slices.Equal[[]string]) +} + +type ServerConn struct { + w http.ResponseWriter + request *http.Request + proxyRequest *connectip.ProxyRequest + local net.Addr + remote net.Addr + answer sync.Once + mu sync.Mutex + ipConn *connectip.Conn + done chan struct{} + closeOnce sync.Once +} + +func (c *ServerConn) Request() *http.Request { + return c.request +} + +func (c *ServerConn) Accept() (*connectip.Conn, error) { + var err error = errors.New("the request was already answered") + c.answer.Do(func() { + var ipConn *connectip.Conn + ipConn, err = (&connectip.Proxy{}).Proxy(c.w, c.proxyRequest) + c.mu.Lock() + c.ipConn = ipConn + c.mu.Unlock() + }) + if err != nil { + return nil, err + } + return c.ipConn, nil +} + +func (c *ServerConn) Reject(status int, header http.Header) { + c.answer.Do(func() { + for k, vv := range header { + for _, v := range vv { + c.w.Header().Add(k, v) + } + } + c.w.WriteHeader(status) + }) +} + +func (c *ServerConn) tunnel() *connectip.Conn { + c.mu.Lock() + defer c.mu.Unlock() + return c.ipConn +} + +func (c *ServerConn) Read(b []byte) (int, error) { + ipConn := c.tunnel() + if ipConn == nil { + return 0, io.ErrClosedPipe + } + return ipConn.ReadPacket(b) +} + +func (c *ServerConn) Write(b []byte) (int, error) { + ipConn := c.tunnel() + if ipConn == nil { + return 0, io.ErrClosedPipe + } + icmp, err := ipConn.WritePacket(b) + if err != nil { + return 0, err + } + if len(icmp) > 0 { + return 0, &PacketTooBigError{ICMP: icmp} + } + return len(b), nil +} + +func (c *ServerConn) Close() error { + c.closeOnce.Do(func() { + c.Reject(http.StatusInternalServerError, nil) + if ipConn := c.tunnel(); ipConn != nil { + ipConn.Close() + } + close(c.done) + }) + return nil +} + +func (c *ServerConn) LocalAddr() net.Addr { + return c.local +} + +func (c *ServerConn) RemoteAddr() net.Addr { + return c.remote +} + +func (c *ServerConn) SetDeadline(time.Time) error { + return nil +} + +func (c *ServerConn) SetReadDeadline(time.Time) error { + return nil +} + +func (c *ServerConn) SetWriteDeadline(time.Time) error { + return nil +} + +func init() { + common.Must(internet.RegisterTransportListener(protocolName, Listen)) +} diff --git a/transport/internet/masque/hub_test.go b/transport/internet/masque/hub_test.go new file mode 100644 index 000000000..c3069fff3 --- /dev/null +++ b/transport/internet/masque/hub_test.go @@ -0,0 +1,150 @@ +package masque + +import ( + "context" + "io" + "net/http" + "net/http/httptest" + "net/url" + "testing" + "time" + + "github.com/stretchr/testify/require" + "github.com/xtls/xray-core/transport/internet/stat" + "github.com/xtls/xray-core/transport/internet/tls" +) + +func TestServerVersions(t *testing.T) { + for _, c := range []struct { + alpn []string + h2, h3 bool + }{ + {alpn: nil, h3: true}, + {alpn: []string{"h3"}, h3: true}, + {alpn: []string{"h2"}, h2: true}, + {alpn: []string{"h2", "http/1.1"}, h2: true}, + {alpn: []string{"h3", "h2"}, h2: true, h3: true}, + {alpn: []string{"http/1.1"}, h3: true}, + } { + h2, h3 := serverVersions(&tls.Config{NextProtocol: c.alpn}) + if h2 != c.h2 || h3 != c.h3 { + t.Errorf("serverVersions(%q) = %v, %v, want %v, %v", c.alpn, h2, h3, c.h2, c.h3) + } + } +} + +func TestPathMatcher(t *testing.T) { + for _, c := range []struct { + path string + request string + want bool + }{ + {DefaultPath, "/.well-known/masque/ip/*/*/", true}, + {DefaultPath, "/.well-known/masque/ip/%2A/%2A/", true}, + {DefaultPath, "/.well-known/masque/ip/*/*", false}, + {DefaultPath, "/.well-known/masque/ip/*/*/?x=1", false}, + {DefaultPath, "/.well-known/masque/ip/192.0.2.1/6/", false}, + {"/masque?target=*&ipproto=*", "/masque?ipproto=*&target=*", true}, + {"/masque?target=*&ipproto=*", "/masque?target=*", false}, + } { + m, err := newPathMatcher(c.path) + require.NoError(t, err) + u, err := url.ParseRequestURI(c.request) + require.NoError(t, err) + if got := m.match(u); got != c.want { + t.Errorf("path %q matching %q = %v, want %v", c.path, c.request, got, c.want) + } + } + _, err := newPathMatcher("masque") + require.Error(t, err) +} + +func connectIPRequest(target string) *http.Request { + r := httptest.NewRequest(http.MethodGet, "https://proxy.example"+target, nil) + r.Method = http.MethodConnect + r.Proto, r.ProtoMajor, r.ProtoMinor = "HTTP/2.0", 2, 0 + r.Header.Set(":protocol", "connect-ip") + r.Header.Set("Capsule-Protocol", "?1") + r.Body = io.NopCloser(&blockingReader{}) + return r +} + +type blockingReader struct{} + +func (*blockingReader) Read([]byte) (int, error) { + time.Sleep(time.Hour) + return 0, io.EOF +} + +func serve(t *testing.T, r *http.Request, handle func(*ServerConn)) *httptest.ResponseRecorder { + t.Helper() + path, err := newPathMatcher(DefaultPath) + require.NoError(t, err) + l := &Listener{path: path, addConn: func(conn stat.Connection) { + go func() { + handle(conn.(*ServerConn)) + conn.Close() + }() + }} + l.ctx, l.cancel = context.WithCancel(context.Background()) + defer l.cancel() + w := httptest.NewRecorder() + done := make(chan struct{}) + go func() { + l.ServeHTTP(w, r) + close(done) + }() + select { + case <-done: + case <-time.After(5 * time.Second): + t.Fatal("ServeHTTP did not return") + } + return w +} + +func TestListenerRejectsOtherRequests(t *testing.T) { + unexpected := func(*ServerConn) { t.Error("an invalid request reached the proxy") } + + require.Equal(t, http.StatusNotFound, serve(t, connectIPRequest("/other"), unexpected).Code) + + get := connectIPRequest(DefaultPath) + get.Method = http.MethodGet + require.Equal(t, http.StatusMethodNotAllowed, serve(t, get, unexpected).Code) + + websocket := connectIPRequest(DefaultPath) + websocket.Header.Set(":protocol", "websocket") + require.Equal(t, http.StatusNotImplemented, serve(t, websocket, unexpected).Code) + + noCapsules := connectIPRequest(DefaultPath) + noCapsules.Header.Del("Capsule-Protocol") + require.Equal(t, http.StatusBadRequest, serve(t, noCapsules, unexpected).Code) +} + +func TestServerConnAnswers(t *testing.T) { + t.Run("reject", func(t *testing.T) { + w := serve(t, connectIPRequest(DefaultPath), func(conn *ServerConn) { + require.Equal(t, "connect-ip", conn.Request().Header.Get(":protocol")) + conn.Reject(http.StatusUnauthorized, http.Header{"WWW-Authenticate": {"Basic"}}) + }) + require.Equal(t, http.StatusUnauthorized, w.Code) + require.Equal(t, "Basic", w.Header().Get("WWW-Authenticate")) + }) + + t.Run("no answer", func(t *testing.T) { + w := serve(t, connectIPRequest(DefaultPath), func(*ServerConn) {}) + require.Equal(t, http.StatusInternalServerError, w.Code) + }) + + t.Run("accept", func(t *testing.T) { + w := serve(t, connectIPRequest(DefaultPath), func(conn *ServerConn) { + ipConn, err := conn.Accept() + require.NoError(t, err) + require.NotNil(t, ipConn) + _, err = conn.Accept() + require.Error(t, err) + conn.Reject(http.StatusForbidden, nil) + }) + require.Equal(t, http.StatusOK, w.Code) + require.Equal(t, "?1", w.Header().Get("Capsule-Protocol")) + }) +}