Proxy: Add MASQUE inbound (IETF CONNECT-IP server, RFC 9484) (#6844)

Completes https://github.com/XTLS/Xray-core/pull/6807 and https://github.com/XTLS/Xray-core/pull/6810
This commit is contained in:
Cluvex
2026-09-27 19:16:54 +00:00
committed by GitHub
parent 65e853ed84
commit 7a018833ec
22 changed files with 3339 additions and 91 deletions
+66
View File
@@ -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
}
+103
View File
@@ -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")
}
}
+20 -1
View File
@@ -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
}
+4
View File
@@ -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,
@@ -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")
}
+108
View File
@@ -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))
}
+49
View File
@@ -0,0 +1,49 @@
package masque
import (
"testing"
"github.com/stretchr/testify/require"
"github.com/xtls/xray-core/common/protocol"
)
func TestValidator(t *testing.T) {
v := newValidator()
user := &protocol.MemoryUser{Email: "U@example.com", Account: &MemoryAccount{Password: "p"}}
require.NoError(t, v.add(user))
for _, u := range []*protocol.MemoryUser{
{Email: "u@example.com", Account: &MemoryAccount{Password: "other"}},
{Account: &MemoryAccount{Password: "p"}},
{Email: "a:b", Account: &MemoryAccount{Password: "p"}},
{Email: "b@example.com", Account: &MemoryAccount{}},
} {
require.Error(t, v.add(u), u.Email)
}
require.Equal(t, user, v.get("u@example.com", "p"))
require.Equal(t, user, v.get("U@EXAMPLE.COM", "p"))
require.Nil(t, v.get("u@example.com", "x"))
require.Nil(t, v.get("x@example.com", "p"))
require.Nil(t, v.get("", ""))
require.Equal(t, user, v.getByEmail("u@example.com"))
require.Equal(t, []*protocol.MemoryUser{user}, v.getAll())
require.Equal(t, int64(1), v.count())
require.True(t, v.contains(user))
removed, err := v.delByEmail("u@EXAMPLE.com")
require.NoError(t, err)
require.Equal(t, user, removed)
_, err = v.delByEmail("u@example.com")
require.Error(t, err)
require.False(t, v.contains(user))
require.Nil(t, v.get("u@example.com", "p"))
require.Zero(t, v.count())
}
func TestAccount(t *testing.T) {
account, err := (&Account{Password: "p"}).AsAccount()
require.NoError(t, err)
require.True(t, account.Equals(&MemoryAccount{Password: "p"}))
require.False(t, account.Equals(&MemoryAccount{Password: "x"}))
require.Equal(t, &Account{Password: "p"}, account.ToProto())
}
+125 -11
View File
@@ -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,
},
+11
View File
@@ -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;
}
+80
View File
@@ -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)
}
+59
View File
@@ -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)
}
}
+550
View File
@@ -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))
}))
}
+328
View File
@@ -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)
}
}
+1 -1
View File
@@ -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
}
+1 -1
View File
@@ -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)
+275 -54
View File
@@ -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)
}
}
}
}
+12 -12
View File
@@ -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()
@@ -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)
+680
View File
@@ -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
}
@@ -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)
}
+422
View File
@@ -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))
}
+150
View File
@@ -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"))
})
}