mirror of
https://github.com/XTLS/Xray-core.git
synced 2026-09-30 13:05:43 +00:00
Completes https://github.com/XTLS/Xray-core/pull/6807 and https://github.com/XTLS/Xray-core/pull/6810
423 lines
11 KiB
Go
423 lines
11 KiB
Go
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))
|
|
}
|