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