Proxy: Add MASQUE outbound & transport (IETF CONNECT-IP, RFC 9484) (#6807)

Closes https://github.com/XTLS/Xray-core/issues/5495#issuecomment-3710683679
This commit is contained in:
Esko Mobius
2026-09-24 02:13:37 +00:00
committed by GitHub
parent 9927942aaa
commit 7d3e44fee2
37 changed files with 5186 additions and 10 deletions
+328
View File
@@ -0,0 +1,328 @@
package masque
import (
"context"
go_errors "errors"
"io"
"net/netip"
"slices"
"sync"
"sync/atomic"
"time"
"golang.zx2c4.com/wireguard/tun"
"github.com/xtls/xray-core/common"
"github.com/xtls/xray-core/common/buf"
"github.com/xtls/xray-core/common/errors"
"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/common/signal"
"github.com/xtls/xray-core/common/task"
"github.com/xtls/xray-core/core"
"github.com/xtls/xray-core/features/policy"
"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/stat"
"github.com/xtls/xray-core/transport/internet/tls"
)
const (
establishTimeout = 10 * time.Second
retryInterval = time.Second
)
type Client struct {
server *protocol.ServerSpec
policyManager policy.Manager
remoteDNS []netip.Addr
ctx context.Context
cancel context.CancelFunc
tunnel atomic.Pointer[tunnel]
mu sync.Mutex
lastErr error
lastErrAt time.Time
}
func NewClient(ctx context.Context, config *ClientConfig) (*Client, error) {
v := core.MustFromContext(ctx)
p := v.GetFeature(policy.ManagerType()).(policy.Manager)
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"`)
}
if config.Server == nil {
return nil, errors.New(`no target server found`)
}
server, err := protocol.NewServerSpecFromPB(config.Server)
if err != nil {
return nil, errors.New("failed to get server spec").Base(err)
}
dns := config.RemoteDns
if len(dns) == 0 {
dns = []string{"1.1.1.1", "1.0.0.1", "2606:4700:4700::1111", "2606:4700:4700::1001"}
}
remoteDNS := make([]netip.Addr, 0, len(dns))
for _, s := range dns {
addr, err := netip.ParseAddr(s)
if err != nil {
return nil, errors.New("invalid remote DNS server ", s).Base(err)
}
remoteDNS = append(remoteDNS, addr)
}
c := &Client{
server: server,
policyManager: p,
remoteDNS: remoteDNS,
}
c.ctx, c.cancel = context.WithCancel(context.Background())
return c, nil
}
func (c *Client) Process(ctx context.Context, link *transport.Link, dialer internet.Dialer) error {
outbounds := session.OutboundsFromContext(ctx)
ob := outbounds[len(outbounds)-1]
if !ob.Target.IsValid() {
return errors.New("target not specified")
}
ob.Name = "masque"
ob.CanSpliceCopy = 3
t, err := c.getTunnel(ctx, dialer)
if err != nil {
return errors.New("failed to establish CONNECT-IP tunnel").Base(err)
}
var newCtx context.Context
var newCancel context.CancelFunc
if session.TimeoutOnlyFromContext(ctx) {
newCtx, newCancel = context.WithCancel(context.Background())
}
sessionPolicy := c.policyManager.ForLevel(0)
ctx, cancel := context.WithCancel(ctx)
timer := signal.CancelAfterInactivity(ctx, func() {
cancel()
if newCancel != nil {
newCancel()
}
}, sessionPolicy.Timeouts.ConnectionIdle)
if newCtx != nil {
ctx = newCtx
}
var reader buf.Reader
var writer buf.Writer
switch ob.Target.Network {
case net.Network_TCP:
var conn net.Conn
var err error
if sessionPolicy.Timeouts.Handshake != 0 {
timeoutCtx, timeoutCancel := context.WithTimeout(ctx, sessionPolicy.Timeouts.Handshake)
conn, err = t.tnet.DialContext(timeoutCtx, "tcp", ob.Target.NetAddr())
timeoutCancel()
} else {
conn, err = t.tnet.Dial("tcp", ob.Target.NetAddr())
}
if err != nil {
return errors.New("failed to create TCP connection").Base(err)
}
defer conn.Close()
reader = buf.NewReader(conn)
writer = buf.NewWriter(conn)
case net.Network_UDP:
conn, err := t.tnet.Dial("udp", ob.Target.NetAddr())
if err != nil {
return errors.New("failed to create UDP connection").Base(err)
}
defer conn.Close()
uc := &wireguard.UDPConnClient{
PacketConn: conn.(*internet.PacketConnWrapper).PacketConn,
Dest: conn.RemoteAddr().(*net.UDPAddr),
}
reader = uc
writer = uc
default:
panic(ob.Target.Network)
}
requestFunc := func() error {
defer timer.SetTimeout(sessionPolicy.Timeouts.DownlinkOnly)
return buf.Copy(link.Reader, writer, buf.UpdateActivity(timer))
}
responseFunc := func() error {
defer timer.SetTimeout(sessionPolicy.Timeouts.UplinkOnly)
return buf.Copy(reader, link.Writer, buf.UpdateActivity(timer))
}
responseDonePost := task.OnSuccess(responseFunc, task.Close(link.Writer))
if err := task.Run(ctx, requestFunc, responseDonePost); err != nil {
common.Interrupt(link.Reader)
common.Interrupt(link.Writer)
return errors.New("connection ends").Base(err)
}
return nil
}
func (c *Client) getTunnel(ctx context.Context, dialer internet.Dialer) (*tunnel, error) {
c.mu.Lock()
defer c.mu.Unlock()
if c.ctx.Err() != nil {
return nil, errors.New("closed")
}
if t := c.tunnel.Load(); t != nil {
select {
case <-t.done:
default:
return t, nil
}
}
if err := ctx.Err(); err != nil {
return nil, err
}
if c.lastErr != nil && time.Since(c.lastErrAt) < retryInterval {
return nil, c.lastErr
}
t, err := c.establish(ctx, dialer)
if err != nil {
c.lastErr, c.lastErrAt = err, time.Now()
return nil, err
}
c.lastErr = nil
c.tunnel.Store(t)
if c.ctx.Err() != nil {
if c.tunnel.CompareAndSwap(t, nil) {
t.close()
}
return nil, errors.New("closed")
}
return t, nil
}
func (c *Client) establish(ctx context.Context, dialer internet.Dialer) (*tunnel, error) {
ctx, cancel := context.WithTimeout(context.WithoutCancel(ctx), establishTimeout)
defer cancel()
defer context.AfterFunc(c.ctx, cancel)()
conn, err := dialer.Dial(ctx, c.server.Destination)
if err != nil {
return nil, err
}
mconn, ok := stat.TryUnwrapStatsConn(conn).(*masque.Conn)
if !ok {
conn.Close()
return nil, errors.New("not a CONNECT-IP connection")
}
t, err := newTunnel(conn, mconn.LocalAddrs(), c.remoteDNS)
if err != nil {
conn.Close()
return nil, err
}
errors.LogInfo(ctx, "MASQUE: tunnel established from ", mconn.LocalAddrs())
return t, nil
}
func (c *Client) Close() error {
c.cancel()
if t := c.tunnel.Swap(nil); t != nil {
t.close()
}
return nil
}
type tunnel struct {
conn stat.Connection
dev tun.Device
tnet *wireguard.Net
done chan struct{}
closeOnce sync.Once
}
func newTunnel(conn stat.Connection, local []netip.Addr, remoteDNS []netip.Addr) (*tunnel, error) {
var dns []netip.Addr
for _, addr := range remoteDNS {
if slices.ContainsFunc(local, func(l netip.Addr) bool { return l.Is4() == addr.Is4() }) {
dns = append(dns, addr)
}
}
if len(dns) == 0 {
errors.LogWarning(context.Background(), "MASQUE: no remote DNS server is reachable from the assigned addresses ", local, ", domain names will fail to resolve")
dns = remoteDNS
}
dev, tnet, _, err := wireguard.CreateNetTUN(local, dns, masque.MinPacketSize, true)
if err != nil {
return nil, err
}
t := &tunnel{
conn: conn,
dev: dev,
tnet: tnet,
done: make(chan struct{}),
}
go t.readFromTunnel()
go t.writeToTunnel()
return t, nil
}
func (t *tunnel) readFromTunnel() {
defer t.close()
b := make([]byte, buf.Size)
for {
n, err := t.conn.Read(b)
if err != nil {
if go_errors.Is(err, io.ErrShortBuffer) {
continue
}
errors.LogInfoInner(context.Background(), err, "MASQUE: tunnel closed")
return
}
t.dev.Write([][]byte{b[:n]}, 0)
}
}
func (t *tunnel) writeToTunnel() {
bufs := [][]byte{make([]byte, masque.MinPacketSize)}
sizes := []int{0}
for {
if _, err := t.dev.Read(bufs, sizes, 0); err != nil {
return
}
if _, err := t.conn.Write(bufs[0][:sizes[0]]); err != nil {
var ptb *masque.PacketTooBigError
if go_errors.As(err, &ptb) {
go t.dev.Write([][]byte{ptb.ICMP}, 0)
}
}
}
}
func (t *tunnel) close() {
t.closeOnce.Do(func() {
close(t.done)
t.conn.Close()
t.dev.Close()
})
}
func init() {
common.Must(common.RegisterConfig((*ClientConfig)(nil), func(ctx context.Context, config interface{}) (interface{}, error) {
return NewClient(ctx, config.(*ClientConfig))
}))
}
+136
View File
@@ -0,0 +1,136 @@
// Code generated by protoc-gen-go. DO NOT EDIT.
// versions:
// protoc-gen-go v1.36.11
// protoc v6.33.5
// source: proxy/masque/config.proto
package masque
import (
protocol "github.com/xtls/xray-core/common/protocol"
protoreflect "google.golang.org/protobuf/reflect/protoreflect"
protoimpl "google.golang.org/protobuf/runtime/protoimpl"
reflect "reflect"
sync "sync"
unsafe "unsafe"
)
const (
// Verify that this generated code is sufficiently up-to-date.
_ = protoimpl.EnforceVersion(20 - protoimpl.MinVersion)
// Verify that runtime/protoimpl is sufficiently up-to-date.
_ = protoimpl.EnforceVersion(protoimpl.MaxVersion - 20)
)
type ClientConfig struct {
state protoimpl.MessageState `protogen:"open.v1"`
Server *protocol.ServerEndpoint `protobuf:"bytes,1,opt,name=server,proto3" json:"server,omitempty"`
RemoteDns []string `protobuf:"bytes,2,rep,name=remote_dns,json=remoteDns,proto3" json:"remote_dns,omitempty"`
unknownFields protoimpl.UnknownFields
sizeCache protoimpl.SizeCache
}
func (x *ClientConfig) Reset() {
*x = ClientConfig{}
mi := &file_proxy_masque_config_proto_msgTypes[0]
ms := protoimpl.X.MessageStateOf(protoimpl.Pointer(x))
ms.StoreMessageInfo(mi)
}
func (x *ClientConfig) String() string {
return protoimpl.X.MessageStringOf(x)
}
func (*ClientConfig) ProtoMessage() {}
func (x *ClientConfig) ProtoReflect() protoreflect.Message {
mi := &file_proxy_masque_config_proto_msgTypes[0]
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 ClientConfig.ProtoReflect.Descriptor instead.
func (*ClientConfig) Descriptor() ([]byte, []int) {
return file_proxy_masque_config_proto_rawDescGZIP(), []int{0}
}
func (x *ClientConfig) GetServer() *protocol.ServerEndpoint {
if x != nil {
return x.Server
}
return nil
}
func (x *ClientConfig) GetRemoteDns() []string {
if x != nil {
return x.RemoteDns
}
return nil
}
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" +
"\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" +
"\x15com.xray.proxy.masqueP\x01Z&github.com/xtls/xray-core/proxy/masque\xaa\x02\x11Xray.Proxy.Masqueb\x06proto3"
var (
file_proxy_masque_config_proto_rawDescOnce sync.Once
file_proxy_masque_config_proto_rawDescData []byte
)
func file_proxy_masque_config_proto_rawDescGZIP() []byte {
file_proxy_masque_config_proto_rawDescOnce.Do(func() {
file_proxy_masque_config_proto_rawDescData = protoimpl.X.CompressGZIP(unsafe.Slice(unsafe.StringData(file_proxy_masque_config_proto_rawDesc), len(file_proxy_masque_config_proto_rawDesc)))
})
return file_proxy_masque_config_proto_rawDescData
}
var file_proxy_masque_config_proto_msgTypes = make([]protoimpl.MessageInfo, 1)
var file_proxy_masque_config_proto_goTypes = []any{
(*ClientConfig)(nil), // 0: xray.proxy.masque.ClientConfig
(*protocol.ServerEndpoint)(nil), // 1: xray.common.protocol.ServerEndpoint
}
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
}
func init() { file_proxy_masque_config_proto_init() }
func file_proxy_masque_config_proto_init() {
if File_proxy_masque_config_proto != nil {
return
}
type x struct{}
out := protoimpl.TypeBuilder{
File: protoimpl.DescBuilder{
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,
NumExtensions: 0,
NumServices: 0,
},
GoTypes: file_proxy_masque_config_proto_goTypes,
DependencyIndexes: file_proxy_masque_config_proto_depIdxs,
MessageInfos: file_proxy_masque_config_proto_msgTypes,
}.Build()
File_proxy_masque_config_proto = out.File
file_proxy_masque_config_proto_goTypes = nil
file_proxy_masque_config_proto_depIdxs = nil
}
+14
View File
@@ -0,0 +1,14 @@
syntax = "proto3";
package xray.proxy.masque;
option csharp_namespace = "Xray.Proxy.Masque";
option go_package = "github.com/xtls/xray-core/proxy/masque";
option java_package = "com.xray.proxy.masque";
option java_multiple_files = true;
import "common/protocol/server_spec.proto";
message ClientConfig {
xray.common.protocol.ServerEndpoint server = 1;
repeated string remote_dns = 2;
}
+17 -10
View File
@@ -199,9 +199,9 @@ func (h *Handler) Process(ctx context.Context, link *transport.Link, dialer inte
return errors.New("failed to create UDP connection").Base(err)
}
defer conn.Close()
c := &udpConnClient{
c := &UDPConnClient{
PacketConn: conn.(*internet.PacketConnWrapper).PacketConn,
dest: conn.RemoteAddr().(*net.UDPAddr),
Dest: conn.RemoteAddr().(*net.UDPAddr),
}
reader = c
writer = c
@@ -378,12 +378,12 @@ func (h *Handler) resolveLocal(host string) (net.IP, error) {
return got[dice.Roll(len(got))], nil
}
type udpConnClient struct {
type UDPConnClient struct {
net.PacketConn
dest *net.UDPAddr
Dest *net.UDPAddr
}
func (c *udpConnClient) ReadMultiBuffer() (buf.MultiBuffer, error) {
func (c *UDPConnClient) ReadMultiBuffer() (buf.MultiBuffer, error) {
b := buf.New()
b.Resize(0, buf.Size)
n, addr, err := c.PacketConn.ReadFrom(b.Bytes())
@@ -402,9 +402,9 @@ func (c *udpConnClient) ReadMultiBuffer() (buf.MultiBuffer, error) {
return buf.MultiBuffer{b}, nil
}
func (c *udpConnClient) WriteMultiBuffer(mb buf.MultiBuffer) error {
func (c *UDPConnClient) WriteMultiBuffer(mb buf.MultiBuffer) error {
for i, b := range mb {
dst := c.dest
dst := c.Dest
if b.UDP != nil {
if b.UDP.Address.Family().IsDomain() {
if b.UDP.Port != net.Port(dst.Port) {
@@ -462,20 +462,27 @@ func (c *cache) run() {
return
}
c.running = true
c.m = make(map[string]entry)
if c.m == nil {
c.m = make(map[string]entry)
}
go c.gc()
}
func (c *cache) gc() {
ticker := time.NewTicker(time.Minute)
for {
now := <-ticker.C
defer ticker.Stop()
for now := range ticker.C {
c.mu.Lock()
for key, entry := range c.m {
if now.After(entry.deadline) {
delete(c.m, key)
}
}
if len(c.m) == 0 {
c.running = false
c.mu.Unlock()
return
}
c.mu.Unlock()
}
}