mirror of
https://github.com/XTLS/Xray-core.git
synced 2026-10-01 13:35:53 +00:00
Compare commits
11
Commits
| Author | SHA1 | Date | |
|---|---|---|---|
|
|
21c4fb34cf | ||
|
|
05921d1755 | ||
|
|
d562d8947d | ||
|
|
dbb1ea30ba | ||
|
|
efc9e6da62 | ||
|
|
8267cf953a | ||
|
|
24e6f6d551 | ||
|
|
dcdfc57ccd | ||
|
|
3461c511aa | ||
|
|
c412e77a9b | ||
|
|
ccb69ea5e2 |
@@ -5,7 +5,8 @@ import (
|
||||
)
|
||||
|
||||
type windowsReader struct {
|
||||
bufs []syscall.WSABuf
|
||||
bufs []syscall.WSABuf
|
||||
ready bool
|
||||
}
|
||||
|
||||
func (r *windowsReader) Init(bs []*Buffer) {
|
||||
@@ -15,6 +16,7 @@ func (r *windowsReader) Init(bs []*Buffer) {
|
||||
for _, b := range bs {
|
||||
r.bufs = append(r.bufs, syscall.WSABuf{Len: uint32(Size), Buf: &b.v[0]})
|
||||
}
|
||||
r.ready = false
|
||||
}
|
||||
|
||||
func (r *windowsReader) Clear() {
|
||||
@@ -25,6 +27,14 @@ func (r *windowsReader) Clear() {
|
||||
}
|
||||
|
||||
func (r *windowsReader) Read(fd uintptr) int32 {
|
||||
// On the first invocation, we return -1 to indicate "not ready"
|
||||
// to make rawConn.Read wait for readability using the runtime's own mechanism
|
||||
// because syscall.WSARecv() is a blocking call when used with nil OVERLAPPED
|
||||
if !r.ready {
|
||||
r.ready = true
|
||||
return -1
|
||||
}
|
||||
|
||||
var nBytes uint32
|
||||
var flags uint32
|
||||
err := syscall.WSARecv(syscall.Handle(fd), &r.bufs[0], uint32(len(r.bufs)), &nBytes, &flags, nil, nil)
|
||||
|
||||
@@ -7,6 +7,7 @@ import (
|
||||
"github.com/xtls/xray-core/common/ctx"
|
||||
"github.com/xtls/xray-core/common/net"
|
||||
"github.com/xtls/xray-core/features/outbound"
|
||||
"github.com/xtls/xray-core/features/routing"
|
||||
)
|
||||
|
||||
//go:linkname IndependentCancelCtx context.newCancelCtx
|
||||
@@ -19,13 +20,14 @@ const (
|
||||
isReverseMuxKey ctx.SessionKey = 4 // is reverse mux
|
||||
sockoptSessionKey ctx.SessionKey = 5 // used by dokodemo to only receive sockopt.Mark
|
||||
trackedConnectionErrorKey ctx.SessionKey = 6 // used by observer to get outbound error
|
||||
timeoutOnlyKey ctx.SessionKey = 7 // mux context's child contexts to only cancel when its own traffic times out
|
||||
allowedNetworkKey ctx.SessionKey = 8 // muxcool server control incoming request tcp/udp
|
||||
fullHandlerKey ctx.SessionKey = 9 // outbound gets full handler
|
||||
mitmAlpn11Key ctx.SessionKey = 10 // used by TLS dialer
|
||||
mitmServerNameKey ctx.SessionKey = 11 // used by TLS dialer
|
||||
dispatcherKey ctx.SessionKey = 7 // used by ss2022 inbounds to get dispatcher
|
||||
timeoutOnlyKey ctx.SessionKey = 8 // mux context's child contexts to only cancel when its own traffic times out
|
||||
allowedNetworkKey ctx.SessionKey = 9 // muxcool server control incoming request tcp/udp
|
||||
fullHandlerKey ctx.SessionKey = 10 // outbound gets full handler
|
||||
mitmAlpn11Key ctx.SessionKey = 11 // used by TLS dialer
|
||||
mitmServerNameKey ctx.SessionKey = 12 // used by TLS dialer
|
||||
|
||||
streamSettingsKey ctx.SessionKey = 12
|
||||
streamSettingsKey ctx.SessionKey = 13
|
||||
)
|
||||
|
||||
func ContextWithInbound(ctx context.Context, inbound *Inbound) context.Context {
|
||||
@@ -127,6 +129,17 @@ func TrackedConnectionError(ctx context.Context, tracker TrackedRequestErrorFeed
|
||||
return context.WithValue(ctx, trackedConnectionErrorKey, tracker)
|
||||
}
|
||||
|
||||
func ContextWithDispatcher(ctx context.Context, dispatcher routing.Dispatcher) context.Context {
|
||||
return context.WithValue(ctx, dispatcherKey, dispatcher)
|
||||
}
|
||||
|
||||
func DispatcherFromContext(ctx context.Context) routing.Dispatcher {
|
||||
if dispatcher, ok := ctx.Value(dispatcherKey).(routing.Dispatcher); ok {
|
||||
return dispatcher
|
||||
}
|
||||
return nil
|
||||
}
|
||||
|
||||
func ContextWithTimeoutOnly(ctx context.Context, only bool) context.Context {
|
||||
return context.WithValue(ctx, timeoutOnlyKey, only)
|
||||
}
|
||||
|
||||
@@ -0,0 +1,53 @@
|
||||
package singbridge
|
||||
|
||||
import (
|
||||
M "github.com/sagernet/sing/common/metadata"
|
||||
N "github.com/sagernet/sing/common/network"
|
||||
"github.com/xtls/xray-core/common/errors"
|
||||
"github.com/xtls/xray-core/common/net"
|
||||
)
|
||||
|
||||
func ToNetwork(network string) net.Network {
|
||||
switch N.NetworkName(network) {
|
||||
case N.NetworkTCP:
|
||||
return net.Network_TCP
|
||||
case N.NetworkUDP:
|
||||
return net.Network_UDP
|
||||
default:
|
||||
return net.Network_Unknown
|
||||
}
|
||||
}
|
||||
|
||||
func ToDestination(socksaddr M.Socksaddr, network net.Network) (net.Destination, error) {
|
||||
// IsFqdn() implicitly checks if the domain name is valid
|
||||
if socksaddr.IsFqdn() {
|
||||
return net.Destination{
|
||||
Network: network,
|
||||
Address: net.DomainAddress(socksaddr.Fqdn),
|
||||
Port: net.Port(socksaddr.Port),
|
||||
}, nil
|
||||
}
|
||||
|
||||
// IsIP() implicitly checks if the IP address is valid
|
||||
if socksaddr.IsIP() {
|
||||
return net.Destination{
|
||||
Network: network,
|
||||
Address: net.IPAddress(socksaddr.Addr.AsSlice()),
|
||||
Port: net.Port(socksaddr.Port),
|
||||
}, nil
|
||||
}
|
||||
|
||||
return net.Destination{}, errors.New("invalid socks address: ", socksaddr)
|
||||
}
|
||||
|
||||
func ToSocksaddr(destination net.Destination) M.Socksaddr {
|
||||
var addr M.Socksaddr
|
||||
switch destination.Address.Family() {
|
||||
case net.AddressFamilyDomain:
|
||||
addr.Fqdn = destination.Address.Domain()
|
||||
default:
|
||||
addr.Addr = M.AddrFromIP(destination.Address.IP())
|
||||
}
|
||||
addr.Port = uint16(destination.Port)
|
||||
return addr
|
||||
}
|
||||
@@ -0,0 +1,72 @@
|
||||
package singbridge
|
||||
|
||||
import (
|
||||
"context"
|
||||
"os"
|
||||
|
||||
M "github.com/sagernet/sing/common/metadata"
|
||||
N "github.com/sagernet/sing/common/network"
|
||||
"github.com/xtls/xray-core/common/net"
|
||||
"github.com/xtls/xray-core/common/net/cnc"
|
||||
"github.com/xtls/xray-core/common/session"
|
||||
"github.com/xtls/xray-core/proxy"
|
||||
"github.com/xtls/xray-core/transport"
|
||||
"github.com/xtls/xray-core/transport/internet"
|
||||
"github.com/xtls/xray-core/transport/pipe"
|
||||
)
|
||||
|
||||
var _ N.Dialer = (*XrayDialer)(nil)
|
||||
|
||||
type XrayDialer struct {
|
||||
internet.Dialer
|
||||
}
|
||||
|
||||
func NewDialer(dialer internet.Dialer) *XrayDialer {
|
||||
return &XrayDialer{dialer}
|
||||
}
|
||||
|
||||
func (d *XrayDialer) DialContext(ctx context.Context, network string, destination M.Socksaddr) (net.Conn, error) {
|
||||
dest, err := ToDestination(destination, ToNetwork(network))
|
||||
if err != nil {
|
||||
return nil, err
|
||||
}
|
||||
return d.Dialer.Dial(ctx, dest)
|
||||
}
|
||||
|
||||
func (d *XrayDialer) ListenPacket(ctx context.Context, destination M.Socksaddr) (net.PacketConn, error) {
|
||||
return nil, os.ErrInvalid
|
||||
}
|
||||
|
||||
type XrayOutboundDialer struct {
|
||||
outbound proxy.Outbound
|
||||
dialer internet.Dialer
|
||||
}
|
||||
|
||||
func NewOutboundDialer(outbound proxy.Outbound, dialer internet.Dialer) *XrayOutboundDialer {
|
||||
return &XrayOutboundDialer{outbound, dialer}
|
||||
}
|
||||
|
||||
func (d *XrayOutboundDialer) DialContext(ctx context.Context, network string, destination M.Socksaddr) (net.Conn, error) {
|
||||
dest, err := ToDestination(destination, ToNetwork(network))
|
||||
if err != nil {
|
||||
return nil, err
|
||||
}
|
||||
outbounds := session.OutboundsFromContext(ctx)
|
||||
if len(outbounds) == 0 {
|
||||
outbounds = []*session.Outbound{{}}
|
||||
ctx = session.ContextWithOutbounds(ctx, outbounds)
|
||||
}
|
||||
ob := outbounds[len(outbounds)-1]
|
||||
ob.Target = dest
|
||||
|
||||
opts := []pipe.Option{pipe.WithSizeLimit(64 * 1024)}
|
||||
uplinkReader, uplinkWriter := pipe.New(opts...)
|
||||
downlinkReader, downlinkWriter := pipe.New(opts...)
|
||||
conn := cnc.NewConnection(cnc.ConnectionInputMulti(downlinkWriter), cnc.ConnectionOutputMulti(uplinkReader))
|
||||
go d.outbound.Process(ctx, &transport.Link{Reader: downlinkReader, Writer: uplinkWriter}, d.dialer)
|
||||
return conn, nil
|
||||
}
|
||||
|
||||
func (d *XrayOutboundDialer) ListenPacket(ctx context.Context, destination M.Socksaddr) (net.PacketConn, error) {
|
||||
return nil, os.ErrInvalid
|
||||
}
|
||||
@@ -0,0 +1,10 @@
|
||||
package singbridge
|
||||
|
||||
import E "github.com/sagernet/sing/common/exceptions"
|
||||
|
||||
func ReturnError(err error) error {
|
||||
if E.IsClosedOrCanceled(err) {
|
||||
return nil
|
||||
}
|
||||
return err
|
||||
}
|
||||
@@ -0,0 +1,58 @@
|
||||
package singbridge
|
||||
|
||||
import (
|
||||
"context"
|
||||
"io"
|
||||
|
||||
M "github.com/sagernet/sing/common/metadata"
|
||||
N "github.com/sagernet/sing/common/network"
|
||||
"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/features/routing"
|
||||
"github.com/xtls/xray-core/transport"
|
||||
)
|
||||
|
||||
var (
|
||||
_ N.TCPConnectionHandler = (*Dispatcher)(nil)
|
||||
_ N.UDPConnectionHandler = (*Dispatcher)(nil)
|
||||
)
|
||||
|
||||
type Dispatcher struct {
|
||||
upstream routing.Dispatcher
|
||||
newErrorFunc func(values ...any) *errors.Error
|
||||
}
|
||||
|
||||
func NewDispatcher(dispatcher routing.Dispatcher, newErrorFunc func(values ...any) *errors.Error) *Dispatcher {
|
||||
return &Dispatcher{
|
||||
upstream: dispatcher,
|
||||
newErrorFunc: newErrorFunc,
|
||||
}
|
||||
}
|
||||
|
||||
func (d *Dispatcher) NewConnection(ctx context.Context, conn net.Conn, metadata M.Metadata) error {
|
||||
dest, err := ToDestination(metadata.Destination, net.Network_TCP)
|
||||
if err != nil {
|
||||
return err
|
||||
}
|
||||
xConn := NewConn(conn)
|
||||
return d.upstream.DispatchLink(ctx, dest, &transport.Link{
|
||||
Reader: xConn,
|
||||
Writer: xConn,
|
||||
})
|
||||
}
|
||||
|
||||
func (d *Dispatcher) NewPacketConnection(ctx context.Context, conn N.PacketConn, metadata M.Metadata) error {
|
||||
dest, err := ToDestination(metadata.Destination, net.Network_UDP)
|
||||
if err != nil {
|
||||
return err
|
||||
}
|
||||
return d.upstream.DispatchLink(ctx, dest, &transport.Link{
|
||||
Reader: buf.NewPacketReader(conn.(io.Reader)),
|
||||
Writer: buf.NewWriter(conn.(io.Writer)),
|
||||
})
|
||||
}
|
||||
|
||||
func (d *Dispatcher) NewError(ctx context.Context, err error) {
|
||||
errors.LogInfo(ctx, err.Error())
|
||||
}
|
||||
@@ -0,0 +1,70 @@
|
||||
package singbridge
|
||||
|
||||
import (
|
||||
"context"
|
||||
|
||||
"github.com/sagernet/sing/common/logger"
|
||||
"github.com/xtls/xray-core/common/errors"
|
||||
)
|
||||
|
||||
var _ logger.ContextLogger = (*XrayLogger)(nil)
|
||||
|
||||
type XrayLogger struct {
|
||||
newError func(values ...any) *errors.Error
|
||||
}
|
||||
|
||||
func NewLogger(newErrorFunc func(values ...any) *errors.Error) *XrayLogger {
|
||||
return &XrayLogger{
|
||||
newErrorFunc,
|
||||
}
|
||||
}
|
||||
|
||||
func (l *XrayLogger) Trace(args ...any) {
|
||||
}
|
||||
|
||||
func (l *XrayLogger) Debug(args ...any) {
|
||||
errors.LogDebug(context.Background(), args...)
|
||||
}
|
||||
|
||||
func (l *XrayLogger) Info(args ...any) {
|
||||
errors.LogInfo(context.Background(), args...)
|
||||
}
|
||||
|
||||
func (l *XrayLogger) Warn(args ...any) {
|
||||
errors.LogWarning(context.Background(), args...)
|
||||
}
|
||||
|
||||
func (l *XrayLogger) Error(args ...any) {
|
||||
errors.LogError(context.Background(), args...)
|
||||
}
|
||||
|
||||
func (l *XrayLogger) Fatal(args ...any) {
|
||||
}
|
||||
|
||||
func (l *XrayLogger) Panic(args ...any) {
|
||||
}
|
||||
|
||||
func (l *XrayLogger) TraceContext(ctx context.Context, args ...any) {
|
||||
}
|
||||
|
||||
func (l *XrayLogger) DebugContext(ctx context.Context, args ...any) {
|
||||
errors.LogDebug(ctx, args...)
|
||||
}
|
||||
|
||||
func (l *XrayLogger) InfoContext(ctx context.Context, args ...any) {
|
||||
errors.LogInfo(ctx, args...)
|
||||
}
|
||||
|
||||
func (l *XrayLogger) WarnContext(ctx context.Context, args ...any) {
|
||||
errors.LogWarning(ctx, args...)
|
||||
}
|
||||
|
||||
func (l *XrayLogger) ErrorContext(ctx context.Context, args ...any) {
|
||||
errors.LogError(ctx, args...)
|
||||
}
|
||||
|
||||
func (l *XrayLogger) FatalContext(ctx context.Context, args ...any) {
|
||||
}
|
||||
|
||||
func (l *XrayLogger) PanicContext(ctx context.Context, args ...any) {
|
||||
}
|
||||
@@ -0,0 +1,107 @@
|
||||
package singbridge
|
||||
|
||||
import (
|
||||
"context"
|
||||
"time"
|
||||
|
||||
B "github.com/sagernet/sing/common/buf"
|
||||
"github.com/sagernet/sing/common/bufio"
|
||||
M "github.com/sagernet/sing/common/metadata"
|
||||
"github.com/xtls/xray-core/common"
|
||||
"github.com/xtls/xray-core/common/buf"
|
||||
"github.com/xtls/xray-core/common/net"
|
||||
"github.com/xtls/xray-core/common/signal"
|
||||
"github.com/xtls/xray-core/transport"
|
||||
)
|
||||
|
||||
func CopyPacketConn(ctx context.Context, inboundConn net.Conn, link *transport.Link, destination net.Destination, serverConn net.PacketConn) error {
|
||||
cancel := func() {
|
||||
common.Interrupt(link.Reader)
|
||||
common.Interrupt(serverConn)
|
||||
}
|
||||
conn := &PacketConnWrapper{
|
||||
Reader: link.Reader,
|
||||
Writer: link.Writer,
|
||||
Dest: destination,
|
||||
Conn: inboundConn,
|
||||
T: signal.CancelAfterInactivity(ctx, cancel, 300*time.Second),
|
||||
}
|
||||
return ReturnError(bufio.CopyPacketConn(ctx, conn, bufio.NewPacketConn(serverConn)))
|
||||
}
|
||||
|
||||
type PacketConnWrapper struct {
|
||||
buf.Reader
|
||||
buf.Writer
|
||||
net.Conn
|
||||
Dest net.Destination
|
||||
cached buf.MultiBuffer
|
||||
|
||||
// A simple patch to avoid goroutine leak since sing infra cannot awake read block by write err
|
||||
T *signal.ActivityTimer
|
||||
}
|
||||
|
||||
func (w *PacketConnWrapper) ReadPacket(buffer *B.Buffer) (addr M.Socksaddr, err error) {
|
||||
w.T.Update()
|
||||
defer func() {
|
||||
if err != nil {
|
||||
// uplinkonly
|
||||
w.T.SetTimeout(2 * time.Second)
|
||||
}
|
||||
}()
|
||||
if w.cached != nil {
|
||||
mb, bb := buf.SplitFirst(w.cached)
|
||||
if bb == nil {
|
||||
w.cached = nil
|
||||
} else {
|
||||
buffer.Write(bb.Bytes())
|
||||
w.cached = mb
|
||||
var destination net.Destination
|
||||
if bb.UDP != nil {
|
||||
destination = *bb.UDP
|
||||
} else {
|
||||
destination = w.Dest
|
||||
}
|
||||
bb.Release()
|
||||
return ToSocksaddr(destination), nil
|
||||
}
|
||||
}
|
||||
mb, err := w.ReadMultiBuffer()
|
||||
nb, bb := buf.SplitFirst(mb)
|
||||
if bb == nil {
|
||||
return M.Socksaddr{}, nil
|
||||
} else {
|
||||
buffer.Write(bb.Bytes())
|
||||
w.cached = nb
|
||||
var destination net.Destination
|
||||
if bb.UDP != nil {
|
||||
destination = *bb.UDP
|
||||
} else {
|
||||
destination = w.Dest
|
||||
}
|
||||
bb.Release()
|
||||
return ToSocksaddr(destination), nil
|
||||
}
|
||||
}
|
||||
|
||||
func (w *PacketConnWrapper) WritePacket(buffer *B.Buffer, destination M.Socksaddr) (err error) {
|
||||
w.T.Update()
|
||||
defer func() {
|
||||
if err != nil {
|
||||
// downlinkonly
|
||||
w.T.SetTimeout(5 * time.Second)
|
||||
}
|
||||
}()
|
||||
endpoint, err := ToDestination(destination, net.Network_UDP)
|
||||
if err != nil {
|
||||
return err
|
||||
}
|
||||
vBuf := buf.New()
|
||||
vBuf.Write(buffer.Bytes())
|
||||
vBuf.UDP = &endpoint
|
||||
return w.WriteMultiBuffer(buf.MultiBuffer{vBuf})
|
||||
}
|
||||
|
||||
func (w *PacketConnWrapper) Close() error {
|
||||
buf.ReleaseMulti(w.cached)
|
||||
return nil
|
||||
}
|
||||
@@ -0,0 +1,81 @@
|
||||
package singbridge
|
||||
|
||||
import (
|
||||
"context"
|
||||
"io"
|
||||
"net"
|
||||
"time"
|
||||
|
||||
"github.com/sagernet/sing/common/bufio"
|
||||
"github.com/xtls/xray-core/common"
|
||||
"github.com/xtls/xray-core/common/buf"
|
||||
"github.com/xtls/xray-core/common/signal"
|
||||
"github.com/xtls/xray-core/transport"
|
||||
)
|
||||
|
||||
func CopyConn(ctx context.Context, inboundConn net.Conn, link *transport.Link, serverConn net.Conn) error {
|
||||
conn := &PipeConnWrapper{
|
||||
W: link.Writer,
|
||||
Conn: inboundConn,
|
||||
}
|
||||
if ir, ok := link.Reader.(io.Reader); ok {
|
||||
conn.R = ir
|
||||
} else {
|
||||
conn.R = &buf.BufferedReader{Reader: link.Reader}
|
||||
}
|
||||
cancel := func() {
|
||||
common.Interrupt(link.Reader)
|
||||
common.Interrupt(serverConn)
|
||||
}
|
||||
conn.T = signal.CancelAfterInactivity(ctx, cancel, 300*time.Second)
|
||||
return ReturnError(bufio.CopyConn(ctx, conn, serverConn))
|
||||
}
|
||||
|
||||
type PipeConnWrapper struct {
|
||||
R io.Reader
|
||||
W buf.Writer
|
||||
net.Conn
|
||||
|
||||
// A simple patch to avoid goroutine leak since sing infra cannot awake read block by write err
|
||||
T *signal.ActivityTimer
|
||||
}
|
||||
|
||||
func (w *PipeConnWrapper) Close() error {
|
||||
return nil
|
||||
}
|
||||
|
||||
func (w *PipeConnWrapper) Read(b []byte) (n int, err error) {
|
||||
w.T.Update()
|
||||
n, err = w.R.Read(b)
|
||||
if err != nil {
|
||||
// uplinkonly
|
||||
w.T.SetTimeout(2 * time.Second)
|
||||
}
|
||||
return
|
||||
}
|
||||
|
||||
func (w *PipeConnWrapper) Write(p []byte) (n int, err error) {
|
||||
w.T.Update()
|
||||
n = len(p)
|
||||
var mb buf.MultiBuffer
|
||||
pLen := len(p)
|
||||
for pLen > 0 {
|
||||
buffer := buf.New()
|
||||
if pLen > buf.Size {
|
||||
_, err = buffer.Write(p[:buf.Size])
|
||||
p = p[buf.Size:]
|
||||
} else {
|
||||
buffer.Write(p)
|
||||
}
|
||||
pLen -= int(buffer.Len())
|
||||
mb = append(mb, buffer)
|
||||
}
|
||||
err = w.W.WriteMultiBuffer(mb)
|
||||
if err != nil {
|
||||
n = 0
|
||||
buf.ReleaseMulti(mb)
|
||||
// downlinkonly
|
||||
w.T.SetTimeout(5 * time.Second)
|
||||
}
|
||||
return
|
||||
}
|
||||
@@ -0,0 +1,66 @@
|
||||
package singbridge
|
||||
|
||||
import (
|
||||
"time"
|
||||
|
||||
"github.com/sagernet/sing/common"
|
||||
"github.com/sagernet/sing/common/bufio"
|
||||
N "github.com/sagernet/sing/common/network"
|
||||
"github.com/xtls/xray-core/common/buf"
|
||||
"github.com/xtls/xray-core/common/net"
|
||||
)
|
||||
|
||||
var (
|
||||
_ buf.Reader = (*Conn)(nil)
|
||||
_ buf.TimeoutReader = (*Conn)(nil)
|
||||
_ buf.Writer = (*Conn)(nil)
|
||||
)
|
||||
|
||||
type Conn struct {
|
||||
net.Conn
|
||||
writer N.VectorisedWriter
|
||||
}
|
||||
|
||||
func NewConn(conn net.Conn) *Conn {
|
||||
writer, _ := bufio.CreateVectorisedWriter(conn)
|
||||
return &Conn{
|
||||
Conn: conn,
|
||||
writer: writer,
|
||||
}
|
||||
}
|
||||
|
||||
func (c *Conn) ReadMultiBuffer() (buf.MultiBuffer, error) {
|
||||
buffer, err := buf.ReadBuffer(c.Conn)
|
||||
if err != nil {
|
||||
return nil, err
|
||||
}
|
||||
return buf.MultiBuffer{buffer}, nil
|
||||
}
|
||||
|
||||
func (c *Conn) ReadMultiBufferTimeout(duration time.Duration) (buf.MultiBuffer, error) {
|
||||
err := c.SetReadDeadline(time.Now().Add(duration))
|
||||
if err != nil {
|
||||
return nil, err
|
||||
}
|
||||
defer c.SetReadDeadline(time.Time{})
|
||||
return c.ReadMultiBuffer()
|
||||
}
|
||||
|
||||
func (c *Conn) WriteMultiBuffer(bufferList buf.MultiBuffer) error {
|
||||
defer buf.ReleaseMulti(bufferList)
|
||||
if c.writer != nil {
|
||||
bytesList := make([][]byte, len(bufferList))
|
||||
for i, buffer := range bufferList {
|
||||
bytesList[i] = buffer.Bytes()
|
||||
}
|
||||
return common.Error(bufio.WriteVectorised(c.writer, bytesList))
|
||||
}
|
||||
// Since this conn is only used by tun, we don't force buffer writes to merge.
|
||||
for _, buffer := range bufferList {
|
||||
_, err := c.Conn.Write(buffer.Bytes())
|
||||
if err != nil {
|
||||
return err
|
||||
}
|
||||
}
|
||||
return nil
|
||||
}
|
||||
@@ -18,15 +18,17 @@ require (
|
||||
github.com/pires/go-proxyproto v0.15.0
|
||||
github.com/refraction-networking/utls v1.8.3-0.20260301010127-aa6edf4b11af
|
||||
github.com/robfig/cron/v3 v3.0.1
|
||||
github.com/sagernet/sing v0.5.1
|
||||
github.com/sagernet/sing-shadowsocks v0.2.7
|
||||
github.com/stretchr/testify v1.12.1
|
||||
github.com/vishvananda/netlink v1.3.1
|
||||
github.com/xtls/reality v0.0.0-20260908062103-8cdf7bf9c7f0
|
||||
go4.org/netipx v0.0.0-20231129151722-fdeea329fbba
|
||||
golang.org/x/crypto v0.55.0
|
||||
golang.org/x/crypto v0.57.0
|
||||
golang.org/x/exp v0.0.0-20240506185415-9bf2ced13842
|
||||
golang.org/x/net v0.58.0
|
||||
golang.org/x/sync v0.22.0
|
||||
golang.org/x/sys v0.47.0
|
||||
golang.org/x/net v0.59.0
|
||||
golang.org/x/sync v0.23.0
|
||||
golang.org/x/sys v0.48.0
|
||||
golang.zx2c4.com/wintun v0.0.0-20230126152724-0fa3db229ce2
|
||||
golang.zx2c4.com/wireguard v0.0.0-20250521234502-f333402bd9cb
|
||||
golang.zx2c4.com/wireguard/windows v1.0.1
|
||||
@@ -55,7 +57,7 @@ require (
|
||||
github.com/vishvananda/netns v0.0.5 // indirect
|
||||
github.com/wlynxg/anet v0.0.5 // indirect
|
||||
go.yaml.in/yaml/v3 v3.0.5 // indirect
|
||||
golang.org/x/text v0.41.0 // indirect
|
||||
golang.org/x/text v0.42.0 // indirect
|
||||
golang.org/x/time v0.14.0 // indirect
|
||||
golang.org/x/tools v0.49.0 // indirect
|
||||
google.golang.org/genproto/googleapis/rpc v0.0.0-20260526163538-3dc84a4a5aaa // indirect
|
||||
|
||||
@@ -76,6 +76,10 @@ github.com/robfig/cron/v3 v3.0.1 h1:WdRxkvbJztn8LMz/QEvLN5sBU+xKpSqwwUO1Pjr4qDs=
|
||||
github.com/robfig/cron/v3 v3.0.1/go.mod h1:eQICP3HwyT7UooqI/z+Ov+PtYAWygg1TEWWzGIFLtro=
|
||||
github.com/rogpeppe/go-internal v1.16.0 h1:O9DK+vNMDVGLr2BeZqmpLeMjiMNkuXfcqntWbZV6S5g=
|
||||
github.com/rogpeppe/go-internal v1.16.0/go.mod h1:DrUVZyrJU+txYW5/1kwtXQSMFio52ZOxX7yM1VHvnxs=
|
||||
github.com/sagernet/sing v0.5.1 h1:mhL/MZVq0TjuvHcpYcFtmSD1BFOxZ/+8ofbNZcg1k1Y=
|
||||
github.com/sagernet/sing v0.5.1/go.mod h1:ARkL0gM13/Iv5VCZmci/NuoOlePoIsW0m7BWfln/Hak=
|
||||
github.com/sagernet/sing-shadowsocks v0.2.7 h1:zaopR1tbHEw5Nk6FAkM05wCslV6ahVegEZaKMv9ipx8=
|
||||
github.com/sagernet/sing-shadowsocks v0.2.7/go.mod h1:0rIKJZBR65Qi0zwdKezt4s57y/Tl1ofkaq6NlkzVuyE=
|
||||
github.com/stretchr/testify v1.12.1 h1:EuwCh5fleGS7H32xRwO3wRGT7DxrDhLAT6FF8MpWDWE=
|
||||
github.com/stretchr/testify v1.12.1/go.mod h1:MDEgiDPPsNp5cuIrHPPCyornHKgEVbtFUmoNlxoYthg=
|
||||
github.com/vishvananda/netlink v1.3.1 h1:3AEMt62VKqz90r0tmNhog0r/PpWKmrEShJU0wJW6bV0=
|
||||
@@ -107,8 +111,8 @@ go4.org/netipx v0.0.0-20231129151722-fdeea329fbba h1:0b9z3AuHCjxk0x/opv64kcgZLBs
|
||||
go4.org/netipx v0.0.0-20231129151722-fdeea329fbba/go.mod h1:PLyyIXexvUFg3Owu6p/WfdlivPbZJsZdgWZlrGope/Y=
|
||||
golang.org/x/crypto v0.0.0-20190308221718-c2843e01d9a2/go.mod h1:djNgcEr1/C05ACkg1iLfiJU5Ep61QUkGW8qpdssI0+w=
|
||||
golang.org/x/crypto v0.0.0-20191011191535-87dc89f01550/go.mod h1:yigFU9vqHzYiE8UmvKecakEJjdnWj3jj499lnFckfCI=
|
||||
golang.org/x/crypto v0.55.0 h1:+KWHjbgOaAQ66dh/YlkZKHlz9ZUlq61AFirAR9ntP8M=
|
||||
golang.org/x/crypto v0.55.0/go.mod h1:uq0V9dE/fzQuJtbnL+2EhWOE63vo164FY8xqEnV9xis=
|
||||
golang.org/x/crypto v0.57.0 h1:3ZVCjf8Ggz7zneR/EHRVx68Ctf+2pmIMP2UFhh9cC6M=
|
||||
golang.org/x/crypto v0.57.0/go.mod h1:Fdz0i5U6CoizGwLda9DttjSk6qlZo25zYNtR+ycvuZA=
|
||||
golang.org/x/exp v0.0.0-20240506185415-9bf2ced13842 h1:vr/HnozRka3pE4EsMEg1lgkXJkTFJCVUX+S/ZT6wYzM=
|
||||
golang.org/x/exp v0.0.0-20240506185415-9bf2ced13842/go.mod h1:XtvwrStGgqGPLc4cjQfWqZHG1YFdYs6swckp8vpsjnc=
|
||||
golang.org/x/lint v0.0.0-20200302205851-738671d3881b/go.mod h1:3xt1FjdF8hUf6vQPIChWIBhFzV8gjjsPE/fR3IyQdNY=
|
||||
@@ -117,12 +121,12 @@ golang.org/x/mod v0.5.1/go.mod h1:5OXOZSfqPIIbmVBIIKWRFfZjPR0E5r58TLhUjH0a2Ro=
|
||||
golang.org/x/net v0.0.0-20190404232315-eb5bcb51f2a3/go.mod h1:t9HGtf8HONx5eT2rtn7q6eTqICYqUVnKs3thJo3Qplg=
|
||||
golang.org/x/net v0.0.0-20190620200207-3b0461eec859/go.mod h1:z5CRVTTTmAJ677TzLLGU+0bjPO0LkuOLi4/5GtJWs/s=
|
||||
golang.org/x/net v0.0.0-20211015210444-4f30a5c0130f/go.mod h1:9nx3DQGgdP8bBQD5qxJ1jj9UTztislL4KSBs9R2vV5Y=
|
||||
golang.org/x/net v0.58.0 h1:ynWG7rqYi4ccpTEuPZ2QGWHktVEM9DMCj9yzDE0Q7To=
|
||||
golang.org/x/net v0.58.0/go.mod h1:YwCddHnFlT7eLQqVprV19OnhLGtc5xOKgE0RyqgfWAU=
|
||||
golang.org/x/net v0.59.0 h1:5zfYln+w5XCxwrnMMJPufRgNoXEaGxl0wo5GqPXyues=
|
||||
golang.org/x/net v0.59.0/go.mod h1:2DA/G1UfVbCpQPeWTmMPGY7Cs2PkBkwu743bVX5PIVg=
|
||||
golang.org/x/sync v0.0.0-20190423024810-112230192c58/go.mod h1:RxMgew5VJxzue5/jJTE5uejpjVlOe/izrB70Jof72aM=
|
||||
golang.org/x/sync v0.0.0-20210220032951-036812b2e83c/go.mod h1:RxMgew5VJxzue5/jJTE5uejpjVlOe/izrB70Jof72aM=
|
||||
golang.org/x/sync v0.22.0 h1:SZjpbeLmrCk4xhRSZFNZW5gFUeCeFgjekvI/+gfScek=
|
||||
golang.org/x/sync v0.22.0/go.mod h1:9xrNwdLfx4jkKbNva9FpL6vEN7evnE43NNNJQ2LF3+0=
|
||||
golang.org/x/sync v0.23.0 h1:KameEIfc1IkluZyXWLn39Wd4tURc6GbCiISGiZm2bQk=
|
||||
golang.org/x/sync v0.23.0/go.mod h1:sUUOizhqBxiL6pEWpqNLUiaJn1ShEbZ6BBqskPbjZm0=
|
||||
golang.org/x/sys v0.0.0-20190215142949-d0b11bdaac8a/go.mod h1:STP8DvDyc/dI5b8T5hshtkjS+E42TnysNCUPdjciGhY=
|
||||
golang.org/x/sys v0.0.0-20190412213103-97732733099d/go.mod h1:h1NjWce9XRLGQEsW7wpKNCjG9DtNlClVuFLEZdDNbEs=
|
||||
golang.org/x/sys v0.0.0-20201119102817-f84b799fce68/go.mod h1:h1NjWce9XRLGQEsW7wpKNCjG9DtNlClVuFLEZdDNbEs=
|
||||
@@ -130,14 +134,14 @@ golang.org/x/sys v0.0.0-20210423082822-04245dca01da/go.mod h1:h1NjWce9XRLGQEsW7w
|
||||
golang.org/x/sys v0.0.0-20211019181941-9d821ace8654/go.mod h1:oPkhp1MJrh7nUepCBck5+mAzfO9JrbApNNgaTdGDITg=
|
||||
golang.org/x/sys v0.2.0/go.mod h1:oPkhp1MJrh7nUepCBck5+mAzfO9JrbApNNgaTdGDITg=
|
||||
golang.org/x/sys v0.10.0/go.mod h1:oPkhp1MJrh7nUepCBck5+mAzfO9JrbApNNgaTdGDITg=
|
||||
golang.org/x/sys v0.47.0 h1:o7XGOvZQCADBQQ4Y7VNq2dRWQR7JmOUW8Kxx4ZsNgWs=
|
||||
golang.org/x/sys v0.47.0/go.mod h1:4GL1E5IUh+htKOUEOaiffhrAeqysfVGipDYzABqnCmw=
|
||||
golang.org/x/sys v0.48.0 h1:bbX/i/6MgT9BVLM9RT1thmxL04yeTAhbEz4SyadbXoo=
|
||||
golang.org/x/sys v0.48.0/go.mod h1:hNLxWAXmnKAxqDtdwIYC4bM9oQPEecfsnNMuSxOs3og=
|
||||
golang.org/x/term v0.0.0-20201126162022-7de9c90e9dd1/go.mod h1:bj7SfCRtBDWHUb9snDiAeCFNEtKQo2Wmx5Cou7ajbmo=
|
||||
golang.org/x/text v0.3.0/go.mod h1:NqM8EUOU14njkJ3fqMW+pc6Ldnwhi/IjpwHt7yyuwOQ=
|
||||
golang.org/x/text v0.3.6/go.mod h1:5Zoc/QRtKVWzQhOtBMvqHzDpF6irO9z98xDceosuGiQ=
|
||||
golang.org/x/text v0.3.7/go.mod h1:u+2+/6zg+i71rQMx5EYifcz6MCKuco9NR6JIITiCfzQ=
|
||||
golang.org/x/text v0.41.0 h1:vz/seA0lnX87Othu2f/0L24RcgrXD9/YFTSuGjj3rH8=
|
||||
golang.org/x/text v0.41.0/go.mod h1:jvf1O8ajNzZqhSrQBPbutR/EB83Cc0CFrezNQIwbb5M=
|
||||
golang.org/x/text v0.42.0 h1:JbOZXgfeCPU9gacVtYliJqOhD+zhrEqK4LfdpmlUZqI=
|
||||
golang.org/x/text v0.42.0/go.mod h1:ojzP1Z+2QtioaF8DTtO8K5q7JWVVYwZKenzujK0Zd0E=
|
||||
golang.org/x/time v0.14.0 h1:MRx4UaLrDotUKUdCIqzPC48t1Y9hANFKIRpNx+Te8PI=
|
||||
golang.org/x/time v0.14.0/go.mod h1:eL/Oa2bBBK0TkX57Fyni+NgnyQQN4LitPmob2Hjnqw4=
|
||||
golang.org/x/tools v0.0.0-20180917221912-90fa682c2a6e/go.mod h1:n7NCudcB/nEzxVGmLbDWY5pfWTLqBcC2KZ6jyYvM4mQ=
|
||||
|
||||
+104
-4
@@ -3,11 +3,14 @@ package conf
|
||||
import (
|
||||
"strings"
|
||||
|
||||
"github.com/sagernet/sing-shadowsocks/shadowaead_2022"
|
||||
C "github.com/sagernet/sing/common"
|
||||
"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/common/task"
|
||||
"github.com/xtls/xray-core/proxy/shadowsocks"
|
||||
"github.com/xtls/xray-core/proxy/shadowsocks_2022"
|
||||
"google.golang.org/protobuf/proto"
|
||||
)
|
||||
|
||||
@@ -27,10 +30,12 @@ func cipherFromString(c string) shadowsocks.CipherType {
|
||||
}
|
||||
|
||||
type ShadowsocksUserConfig struct {
|
||||
Cipher string `json:"method"`
|
||||
Password string `json:"password"`
|
||||
Level byte `json:"level"`
|
||||
Email string `json:"email"`
|
||||
Cipher string `json:"method"`
|
||||
Password string `json:"password"`
|
||||
Level byte `json:"level"`
|
||||
Email string `json:"email"`
|
||||
Address *Address `json:"address"`
|
||||
Port uint16 `json:"port"`
|
||||
}
|
||||
|
||||
type ShadowsocksServerConfig struct {
|
||||
@@ -50,6 +55,10 @@ func (v *ShadowsocksServerConfig) Build() (proto.Message, error) {
|
||||
v.Users = v.Clients
|
||||
}
|
||||
|
||||
if C.Contains(shadowaead_2022.List, v.Cipher) {
|
||||
return buildShadowsocks2022(v)
|
||||
}
|
||||
|
||||
config := new(shadowsocks.ServerConfig)
|
||||
config.Network = v.NetworkList.Build()
|
||||
|
||||
@@ -101,6 +110,72 @@ func (v *ShadowsocksServerConfig) Build() (proto.Message, error) {
|
||||
return config, nil
|
||||
}
|
||||
|
||||
func buildShadowsocks2022(v *ShadowsocksServerConfig) (proto.Message, error) {
|
||||
if len(v.Users) == 0 {
|
||||
config := new(shadowsocks_2022.ServerConfig)
|
||||
config.Method = v.Cipher
|
||||
config.Key = v.Password
|
||||
config.Network = v.NetworkList.Build()
|
||||
config.Email = v.Email
|
||||
return config, nil
|
||||
}
|
||||
|
||||
if v.Cipher == "" {
|
||||
return nil, errors.New("shadowsocks 2022 (multi-user): missing server method")
|
||||
}
|
||||
if !strings.Contains(v.Cipher, "aes") {
|
||||
return nil, errors.New("shadowsocks 2022 (multi-user): only blake3-aes-*-gcm methods are supported")
|
||||
}
|
||||
|
||||
if v.Users[0].Address == nil {
|
||||
config := new(shadowsocks_2022.MultiUserServerConfig)
|
||||
config.Method = v.Cipher
|
||||
config.Key = v.Password
|
||||
config.Network = v.NetworkList.Build()
|
||||
|
||||
config.Users = make([]*protocol.User, len(v.Users))
|
||||
processUser := func(idx int) error {
|
||||
user := v.Users[idx]
|
||||
if user.Cipher != "" {
|
||||
return errors.New("shadowsocks 2022 (multi-user): users must have empty method")
|
||||
}
|
||||
account := &shadowsocks_2022.Account{
|
||||
Key: user.Password,
|
||||
}
|
||||
config.Users[idx] = &protocol.User{
|
||||
Email: user.Email,
|
||||
Level: uint32(user.Level),
|
||||
Account: serial.ToTypedMessage(account),
|
||||
}
|
||||
return nil
|
||||
}
|
||||
if err := task.ParallelForN(len(v.Users), processUser); err != nil {
|
||||
return nil, err
|
||||
}
|
||||
return config, nil
|
||||
}
|
||||
|
||||
config := new(shadowsocks_2022.RelayServerConfig)
|
||||
config.Method = v.Cipher
|
||||
config.Key = v.Password
|
||||
config.Network = v.NetworkList.Build()
|
||||
for _, user := range v.Users {
|
||||
if user.Cipher != "" {
|
||||
return nil, errors.New("shadowsocks 2022 (relay): users must have empty method")
|
||||
}
|
||||
if user.Address == nil {
|
||||
return nil, errors.New("shadowsocks 2022 (relay): all users must have relay address")
|
||||
}
|
||||
config.Destinations = append(config.Destinations, &shadowsocks_2022.RelayDestination{
|
||||
Key: user.Password,
|
||||
Email: user.Email,
|
||||
Address: user.Address.Build(),
|
||||
Port: uint32(user.Port),
|
||||
})
|
||||
}
|
||||
return config, nil
|
||||
}
|
||||
|
||||
type ShadowsocksServerTarget struct {
|
||||
Address *Address `json:"address"`
|
||||
Port uint16 `json:"port"`
|
||||
@@ -139,8 +214,33 @@ func (v *ShadowsocksClientConfig) Build() (proto.Message, error) {
|
||||
return nil, errors.New(`Shadowsocks settings: "servers" should have one and only one member. Multiple endpoints in "servers" should use multiple Shadowsocks outbounds and routing balancer instead`)
|
||||
}
|
||||
|
||||
if len(v.Servers) == 1 {
|
||||
server := v.Servers[0]
|
||||
if C.Contains(shadowaead_2022.List, server.Cipher) {
|
||||
if server.Address == nil {
|
||||
return nil, errors.New("Shadowsocks server address is not set.")
|
||||
}
|
||||
if server.Port == 0 {
|
||||
return nil, errors.New("Invalid Shadowsocks port.")
|
||||
}
|
||||
if server.Password == "" {
|
||||
return nil, errors.New("Shadowsocks password is not specified.")
|
||||
}
|
||||
|
||||
config := new(shadowsocks_2022.ClientConfig)
|
||||
config.Address = server.Address.Build()
|
||||
config.Port = uint32(server.Port)
|
||||
config.Method = server.Cipher
|
||||
config.Key = server.Password
|
||||
return config, nil
|
||||
}
|
||||
}
|
||||
|
||||
config := new(shadowsocks.ClientConfig)
|
||||
for _, server := range v.Servers {
|
||||
if C.Contains(shadowaead_2022.List, server.Cipher) {
|
||||
return nil, errors.New("Shadowsocks 2022 accept no multi servers")
|
||||
}
|
||||
if server.Address == nil {
|
||||
return nil, errors.New("Shadowsocks server address is not set.")
|
||||
}
|
||||
|
||||
@@ -14,7 +14,6 @@ import (
|
||||
googleuuid "github.com/google/uuid"
|
||||
"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/finalmask/fragment"
|
||||
"github.com/xtls/xray-core/transport/internet/finalmask/header/custom"
|
||||
"github.com/xtls/xray-core/transport/internet/finalmask/mkcp/aes128gcm"
|
||||
@@ -909,22 +908,13 @@ func (c *Realm) Build() (proto.Message, error) {
|
||||
}
|
||||
|
||||
type UDPHop struct {
|
||||
Sockopt *SocketConfig `json:"sockopt"`
|
||||
Mode string `json:"mode"`
|
||||
Interval Int32Range `json:"interval"`
|
||||
RemotePorts PortList `json:"remotePorts"`
|
||||
RemoteIPs []string `json:"remoteIPs"`
|
||||
Mode string `json:"mode"`
|
||||
Interval Int32Range `json:"interval"`
|
||||
RemoteIPs []string `json:"remoteIPs"`
|
||||
RemotePorts PortList `json:"remotePorts"`
|
||||
}
|
||||
|
||||
func (c *UDPHop) Build() (proto.Message, error) {
|
||||
var sockopt *internet.SocketConfig
|
||||
if c.Sockopt != nil {
|
||||
var err error
|
||||
sockopt, err = c.Sockopt.Build()
|
||||
if err != nil {
|
||||
return nil, err
|
||||
}
|
||||
}
|
||||
var local, remote, remoteOnce bool
|
||||
for _, mode := range strings.Split(c.Mode, ",") {
|
||||
switch strings.ToLower(mode) {
|
||||
@@ -953,14 +943,13 @@ func (c *UDPHop) Build() (proto.Message, error) {
|
||||
return nil, errors.New("invalid ip ", ip)
|
||||
}
|
||||
return &udphop.Config{
|
||||
Sockopt: sockopt,
|
||||
Local: local,
|
||||
Remote: remote,
|
||||
RemoteOnce: remoteOnce,
|
||||
IntervalMin: int64(c.Interval.From),
|
||||
IntervalMax: int64(c.Interval.To),
|
||||
RemotePorts: c.RemotePorts.Build().Ports(),
|
||||
RemoteIPs: remoteIPs,
|
||||
RemotePorts: c.RemotePorts.Build().Ports(),
|
||||
}, nil
|
||||
}
|
||||
|
||||
|
||||
@@ -36,6 +36,8 @@ func (p TransportProtocol) Build() (string, error) {
|
||||
return "", errors.PrintRemovedFeatureError("QUIC transport (without web service, etc.)", "XHTTP stream-one H3")
|
||||
case "hysteria":
|
||||
return "hysteria", nil
|
||||
case "xdrive":
|
||||
return "xdrive", nil
|
||||
default:
|
||||
return "", errors.New("Config: unknown transport protocol: ", p)
|
||||
}
|
||||
@@ -59,6 +61,7 @@ type StreamConfig struct {
|
||||
WSSettings *WebSocketConfig `json:"wsSettings"`
|
||||
HTTPUPGRADESettings *HttpUpgradeConfig `json:"httpupgradeSettings"`
|
||||
HysteriaSettings *HysteriaConfig `json:"hysteriaSettings"`
|
||||
XDRIVESettings *XDriveConfig `json:"xdriveSettings"`
|
||||
SocketSettings *SocketConfig `json:"sockopt"`
|
||||
}
|
||||
|
||||
@@ -192,6 +195,16 @@ func (c *StreamConfig) Build() (*internet.StreamConfig, error) {
|
||||
Settings: serial.ToTypedMessage(hs),
|
||||
})
|
||||
}
|
||||
if c.XDRIVESettings != nil {
|
||||
xs, err := c.XDRIVESettings.Build()
|
||||
if err != nil {
|
||||
return nil, errors.New("Failed to build XDRIVE config.").Base(err)
|
||||
}
|
||||
config.TransportSettings = append(config.TransportSettings, &internet.TransportConfig{
|
||||
ProtocolName: "xdrive",
|
||||
Settings: serial.ToTypedMessage(xs),
|
||||
})
|
||||
}
|
||||
if c.SocketSettings != nil {
|
||||
ss, err := c.SocketSettings.Build()
|
||||
if err != nil {
|
||||
|
||||
@@ -23,6 +23,7 @@ import (
|
||||
"github.com/xtls/xray-core/transport/internet/splithttp"
|
||||
"github.com/xtls/xray-core/transport/internet/tcp"
|
||||
"github.com/xtls/xray-core/transport/internet/websocket"
|
||||
"github.com/xtls/xray-core/transport/internet/xdrive"
|
||||
"google.golang.org/protobuf/proto"
|
||||
)
|
||||
|
||||
@@ -794,3 +795,50 @@ func readFileOrString(f string, s []string) ([]byte, error) {
|
||||
}
|
||||
return nil, errors.New("both file and bytes are empty.")
|
||||
}
|
||||
|
||||
type XDriveConfig struct {
|
||||
RemoteFolder string `json:"remoteFolder"`
|
||||
Service string `json:"service"`
|
||||
Secrets []string `json:"secrets"`
|
||||
SegmentBytes uint32 `json:"segmentBytes"`
|
||||
FlushIntervalMs uint32 `json:"flushIntervalMs"`
|
||||
PollIntervalMs uint32 `json:"pollIntervalMs"`
|
||||
MaxPollIntervalMs uint32 `json:"maxPollIntervalMs"`
|
||||
SessionTTLSeconds uint32 `json:"sessionTtlSeconds"`
|
||||
Concurrency uint32 `json:"concurrency"`
|
||||
EagerWindowMs uint32 `json:"eagerWindowMs"`
|
||||
HoleTimeoutMs uint32 `json:"holeTimeoutMs"`
|
||||
Template json.RawMessage `json:"template"`
|
||||
}
|
||||
|
||||
// Build implements Buildable.
|
||||
func (c *XDriveConfig) Build() (proto.Message, error) {
|
||||
switch c.Service {
|
||||
case "local":
|
||||
case "Google Drive":
|
||||
if len(c.Secrets) != 3 {
|
||||
return nil, errors.New("Google Drive needs 3 secrets in order of ClientID, ClientSecret, RefreshToken")
|
||||
}
|
||||
case "template":
|
||||
if len(c.Template) == 0 {
|
||||
return nil, errors.New(`service "template" needs a "template" object`)
|
||||
}
|
||||
default:
|
||||
return nil, errors.New("unsupported service")
|
||||
}
|
||||
config := &xdrive.Config{
|
||||
RemoteFolder: c.RemoteFolder,
|
||||
Service: c.Service,
|
||||
Secrets: c.Secrets,
|
||||
SegmentBytes: c.SegmentBytes,
|
||||
FlushIntervalMs: c.FlushIntervalMs,
|
||||
PollIntervalMs: c.PollIntervalMs,
|
||||
MaxPollIntervalMs: c.MaxPollIntervalMs,
|
||||
SessionTtlSeconds: c.SessionTTLSeconds,
|
||||
Concurrency: c.Concurrency,
|
||||
EagerWindowMs: c.EagerWindowMs,
|
||||
HoleTimeoutMs: c.HoleTimeoutMs,
|
||||
Template: string(c.Template),
|
||||
}
|
||||
return config, nil
|
||||
}
|
||||
|
||||
@@ -291,3 +291,76 @@ func TestHeaderCustomUDPBuildRejectsExprWithoutArgs(t *testing.T) {
|
||||
t.Fatalf("expected transform arg rejection, got %v", err)
|
||||
}
|
||||
}
|
||||
|
||||
func TestXDriveStreamConfig(t *testing.T) {
|
||||
config := new(StreamConfig)
|
||||
if err := json.Unmarshal([]byte(`{
|
||||
"method": "xdrive",
|
||||
"xdriveSettings": {
|
||||
"remoteFolder": "/tmp/xdrive",
|
||||
"service": "local"
|
||||
}
|
||||
}`), config); err != nil {
|
||||
t.Fatalf("Unmarshal: %v", err)
|
||||
}
|
||||
|
||||
built, err := config.Build()
|
||||
if err != nil {
|
||||
t.Fatalf("Build: %v", err)
|
||||
}
|
||||
if built.ProtocolName != "xdrive" {
|
||||
t.Fatalf("ProtocolName is %q, want %q", built.ProtocolName, "xdrive")
|
||||
}
|
||||
if len(built.TransportSettings) != 1 || built.TransportSettings[0].ProtocolName != "xdrive" {
|
||||
t.Fatalf("TransportSettings is %v, want a single xdrive entry", built.TransportSettings)
|
||||
}
|
||||
}
|
||||
|
||||
func TestXDriveRejectsUnknownService(t *testing.T) {
|
||||
config := new(XDriveConfig)
|
||||
if err := json.Unmarshal([]byte(`{"remoteFolder": "/tmp/xdrive", "service": "Dropbox"}`), config); err != nil {
|
||||
t.Fatalf("Unmarshal: %v", err)
|
||||
}
|
||||
if _, err := config.Build(); err == nil {
|
||||
t.Fatal("Build accepted an unsupported service")
|
||||
}
|
||||
}
|
||||
|
||||
func TestXDriveTemplateStreamConfig(t *testing.T) {
|
||||
config := new(StreamConfig)
|
||||
if err := json.Unmarshal([]byte(`{
|
||||
"method": "xdrive",
|
||||
"xdriveSettings": {
|
||||
"remoteFolder": "folder",
|
||||
"service": "template",
|
||||
"secrets": ["user", "pass"],
|
||||
"template": {
|
||||
"flatten": true,
|
||||
"auth": {"type": "basic", "username": "{secret0}", "password": "{secret1}"},
|
||||
"put": {"method": "PUT", "url": "https://dav.example/{folder}/{name}"},
|
||||
"get": {"method": "GET", "url": "https://dav.example/{folder}/{name}"},
|
||||
"delete": {"method": "DELETE", "url": "https://dav.example/{folder}/{name}"},
|
||||
"list": {"method": "PROPFIND", "url": "https://dav.example/{folder}/", "namesRegex": "<d:href>/folder/([^<]+)</d:href>"}
|
||||
}
|
||||
}
|
||||
}`), config); err != nil {
|
||||
t.Fatalf("Unmarshal: %v", err)
|
||||
}
|
||||
built, err := config.Build()
|
||||
if err != nil {
|
||||
t.Fatalf("Build: %v", err)
|
||||
}
|
||||
if built.ProtocolName != "xdrive" {
|
||||
t.Fatalf("ProtocolName is %q, want xdrive", built.ProtocolName)
|
||||
}
|
||||
}
|
||||
|
||||
func TestXDriveTemplateNeedsTemplate(t *testing.T) {
|
||||
config := new(XDriveConfig)
|
||||
if err := json.Unmarshal([]byte(`{"remoteFolder": "f", "service": "template"}`), config); err != nil {
|
||||
t.Fatalf("Unmarshal: %v", err)
|
||||
}
|
||||
if _, err := config.Build(); err == nil {
|
||||
t.Fatal("Build accepted a template service without a template")
|
||||
}
|
||||
}
|
||||
|
||||
@@ -20,6 +20,7 @@ type TunConfig struct {
|
||||
UserLevel uint32 `json:"userLevel"`
|
||||
AutoSystemRoutingTable []string `json:"autoSystemRoutingTable"`
|
||||
AutoOutboundsInterface *string `json:"autoOutboundsInterface"`
|
||||
Stack string `json:"stack"`
|
||||
}
|
||||
|
||||
func (v *TunConfig) Build() (proto.Message, error) {
|
||||
@@ -31,6 +32,7 @@ func (v *TunConfig) Build() (proto.Message, error) {
|
||||
DNS: v.DNS,
|
||||
UserLevel: v.UserLevel,
|
||||
AutoSystemRoutingTable: v.AutoSystemRoutingTable,
|
||||
Stack: v.Stack,
|
||||
}
|
||||
if v.AutoOutboundsInterface != nil {
|
||||
config.AutoOutboundsInterface = *v.AutoOutboundsInterface
|
||||
@@ -52,6 +54,11 @@ func (v *TunConfig) Build() (proto.Message, error) {
|
||||
if config.MTU == 0 {
|
||||
config.MTU = 1500
|
||||
}
|
||||
switch config.Stack {
|
||||
case "", "gvisor", "system":
|
||||
default:
|
||||
return nil, fmt.Errorf("unknown tun stack: %s (must be \"gvisor\" or \"system\")", config.Stack)
|
||||
}
|
||||
return config, nil
|
||||
}
|
||||
|
||||
|
||||
+7
-23
@@ -59,14 +59,13 @@ func (c *WireGuardPeerConfig) Build() (*wireguard.PeerConfig, error) {
|
||||
type WireGuardConfig struct {
|
||||
IsClient bool `json:""`
|
||||
|
||||
NoKernelTun bool `json:"noKernelTun"`
|
||||
SecretKey string `json:"secretKey"`
|
||||
Address []string `json:"address"`
|
||||
Peers []*WireGuardPeerConfig `json:"peers"`
|
||||
MTU int32 `json:"mtu"`
|
||||
Reserved []byte `json:"reserved"`
|
||||
DomainStrategy string `json:"domainStrategy"`
|
||||
DNS []string `json:"remoteDNS"`
|
||||
NoKernelTun bool `json:"noKernelTun"`
|
||||
SecretKey string `json:"secretKey"`
|
||||
Address []string `json:"address"`
|
||||
Peers []*WireGuardPeerConfig `json:"peers"`
|
||||
MTU int32 `json:"mtu"`
|
||||
Reserved []byte `json:"reserved"`
|
||||
DNS []string `json:"remoteDNS"`
|
||||
}
|
||||
|
||||
func (c *WireGuardConfig) Build() (proto.Message, error) {
|
||||
@@ -125,21 +124,6 @@ func (c *WireGuardConfig) Build() (proto.Message, error) {
|
||||
}
|
||||
config.Reserved = c.Reserved
|
||||
|
||||
switch strings.ToLower(c.DomainStrategy) {
|
||||
case "forceip", "":
|
||||
config.DomainStrategy = wireguard.DeviceConfig_FORCE_IP
|
||||
case "forceipv4":
|
||||
config.DomainStrategy = wireguard.DeviceConfig_FORCE_IP4
|
||||
case "forceipv6":
|
||||
config.DomainStrategy = wireguard.DeviceConfig_FORCE_IP6
|
||||
case "forceipv4v6":
|
||||
config.DomainStrategy = wireguard.DeviceConfig_FORCE_IP46
|
||||
case "forceipv6v4":
|
||||
config.DomainStrategy = wireguard.DeviceConfig_FORCE_IP64
|
||||
default:
|
||||
return nil, errors.New("unsupported domain strategy: ", c.DomainStrategy)
|
||||
}
|
||||
|
||||
config.IsClient = c.IsClient
|
||||
config.NoKernelTun = c.NoKernelTun
|
||||
config.DNS = c.DNS
|
||||
|
||||
@@ -13,6 +13,7 @@ import (
|
||||
"github.com/xtls/xray-core/infra/conf"
|
||||
"github.com/xtls/xray-core/infra/conf/serial"
|
||||
"github.com/xtls/xray-core/proxy/shadowsocks"
|
||||
"github.com/xtls/xray-core/proxy/shadowsocks_2022"
|
||||
"github.com/xtls/xray-core/proxy/trojan"
|
||||
vlessin "github.com/xtls/xray-core/proxy/vless/inbound"
|
||||
vmessin "github.com/xtls/xray-core/proxy/vmess/inbound"
|
||||
@@ -85,6 +86,8 @@ func extractInboundUsers(inb *core.InboundHandlerConfig) []*protocol.User {
|
||||
return ty.Users
|
||||
case *shadowsocks.ServerConfig:
|
||||
return ty.Users
|
||||
case *shadowsocks_2022.MultiUserServerConfig:
|
||||
return ty.Users
|
||||
default:
|
||||
fmt.Println("unsupported inbound type")
|
||||
}
|
||||
|
||||
@@ -60,6 +60,7 @@ import (
|
||||
_ "github.com/xtls/xray-core/transport/internet/tls"
|
||||
_ "github.com/xtls/xray-core/transport/internet/udp"
|
||||
_ "github.com/xtls/xray-core/transport/internet/websocket"
|
||||
_ "github.com/xtls/xray-core/transport/internet/xdrive"
|
||||
|
||||
// Transport headers
|
||||
_ "github.com/xtls/xray-core/transport/internet/headers/http"
|
||||
|
||||
@@ -4,6 +4,7 @@ import (
|
||||
"context"
|
||||
"crypto/rand"
|
||||
"io"
|
||||
"strings"
|
||||
"sync/atomic"
|
||||
"time"
|
||||
|
||||
@@ -163,9 +164,12 @@ func getDefaultFinalRule(inbound *session.Inbound) *FinalRule {
|
||||
switch inbound.Name {
|
||||
case "vless-reverse":
|
||||
return defaultBlockAllRule
|
||||
case "vless", "vmess", "trojan", "hysteria", "wireguard", "shadowsocks":
|
||||
case "vless", "vmess", "trojan", "hysteria", "wireguard":
|
||||
return defaultBlockPrivateRule
|
||||
default:
|
||||
if strings.HasPrefix(inbound.Name, "shadowsocks") {
|
||||
return defaultBlockPrivateRule
|
||||
}
|
||||
}
|
||||
return nil
|
||||
}
|
||||
|
||||
@@ -236,14 +236,14 @@ type UDPReader struct {
|
||||
|
||||
func (r *UDPReader) ReadFrom(p []byte) (n int, addr *net.Destination, err error) {
|
||||
for {
|
||||
var buf [hysteria.MaxDatagramFrameSize]byte
|
||||
var packet [1500]byte
|
||||
|
||||
n, err := r.reader.Read(buf[:])
|
||||
n, err := r.reader.Read(packet[:])
|
||||
if err != nil {
|
||||
return 0, nil, err
|
||||
}
|
||||
|
||||
msg, err := ParseUDPMessage(buf[:n])
|
||||
msg, err := ParseUDPMessage(packet[:n])
|
||||
if err != nil {
|
||||
continue
|
||||
}
|
||||
|
||||
@@ -0,0 +1,33 @@
|
||||
package shadowsocks_2022
|
||||
|
||||
import (
|
||||
"google.golang.org/protobuf/proto"
|
||||
|
||||
"github.com/xtls/xray-core/common/protocol"
|
||||
)
|
||||
|
||||
// MemoryAccount is an account type converted from Account.
|
||||
type MemoryAccount struct {
|
||||
Key string
|
||||
}
|
||||
|
||||
// AsAccount implements protocol.AsAccount.
|
||||
func (u *Account) AsAccount() (protocol.Account, error) {
|
||||
return &MemoryAccount{
|
||||
Key: u.GetKey(),
|
||||
}, nil
|
||||
}
|
||||
|
||||
// Equals implements protocol.Account.Equals().
|
||||
func (a *MemoryAccount) Equals(another protocol.Account) bool {
|
||||
if account, ok := another.(*MemoryAccount); ok {
|
||||
return a.Key == account.Key
|
||||
}
|
||||
return false
|
||||
}
|
||||
|
||||
func (a *MemoryAccount) ToProto() proto.Message {
|
||||
return &Account{
|
||||
Key: a.Key,
|
||||
}
|
||||
}
|
||||
@@ -0,0 +1,523 @@
|
||||
// Code generated by protoc-gen-go. DO NOT EDIT.
|
||||
// versions:
|
||||
// protoc-gen-go v1.36.11
|
||||
// protoc v6.33.5
|
||||
// source: proxy/shadowsocks_2022/config.proto
|
||||
|
||||
package shadowsocks_2022
|
||||
|
||||
import (
|
||||
net "github.com/xtls/xray-core/common/net"
|
||||
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 ServerConfig struct {
|
||||
state protoimpl.MessageState `protogen:"open.v1"`
|
||||
Method string `protobuf:"bytes,1,opt,name=method,proto3" json:"method,omitempty"`
|
||||
Key string `protobuf:"bytes,2,opt,name=key,proto3" json:"key,omitempty"`
|
||||
Email string `protobuf:"bytes,3,opt,name=email,proto3" json:"email,omitempty"`
|
||||
Level int32 `protobuf:"varint,4,opt,name=level,proto3" json:"level,omitempty"`
|
||||
Network []net.Network `protobuf:"varint,5,rep,packed,name=network,proto3,enum=xray.common.net.Network" json:"network,omitempty"`
|
||||
unknownFields protoimpl.UnknownFields
|
||||
sizeCache protoimpl.SizeCache
|
||||
}
|
||||
|
||||
func (x *ServerConfig) Reset() {
|
||||
*x = ServerConfig{}
|
||||
mi := &file_proxy_shadowsocks_2022_config_proto_msgTypes[0]
|
||||
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_shadowsocks_2022_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 ServerConfig.ProtoReflect.Descriptor instead.
|
||||
func (*ServerConfig) Descriptor() ([]byte, []int) {
|
||||
return file_proxy_shadowsocks_2022_config_proto_rawDescGZIP(), []int{0}
|
||||
}
|
||||
|
||||
func (x *ServerConfig) GetMethod() string {
|
||||
if x != nil {
|
||||
return x.Method
|
||||
}
|
||||
return ""
|
||||
}
|
||||
|
||||
func (x *ServerConfig) GetKey() string {
|
||||
if x != nil {
|
||||
return x.Key
|
||||
}
|
||||
return ""
|
||||
}
|
||||
|
||||
func (x *ServerConfig) GetEmail() string {
|
||||
if x != nil {
|
||||
return x.Email
|
||||
}
|
||||
return ""
|
||||
}
|
||||
|
||||
func (x *ServerConfig) GetLevel() int32 {
|
||||
if x != nil {
|
||||
return x.Level
|
||||
}
|
||||
return 0
|
||||
}
|
||||
|
||||
func (x *ServerConfig) GetNetwork() []net.Network {
|
||||
if x != nil {
|
||||
return x.Network
|
||||
}
|
||||
return nil
|
||||
}
|
||||
|
||||
type MultiUserServerConfig struct {
|
||||
state protoimpl.MessageState `protogen:"open.v1"`
|
||||
Method string `protobuf:"bytes,1,opt,name=method,proto3" json:"method,omitempty"`
|
||||
Key string `protobuf:"bytes,2,opt,name=key,proto3" json:"key,omitempty"`
|
||||
Users []*protocol.User `protobuf:"bytes,3,rep,name=users,proto3" json:"users,omitempty"`
|
||||
Network []net.Network `protobuf:"varint,4,rep,packed,name=network,proto3,enum=xray.common.net.Network" json:"network,omitempty"`
|
||||
unknownFields protoimpl.UnknownFields
|
||||
sizeCache protoimpl.SizeCache
|
||||
}
|
||||
|
||||
func (x *MultiUserServerConfig) Reset() {
|
||||
*x = MultiUserServerConfig{}
|
||||
mi := &file_proxy_shadowsocks_2022_config_proto_msgTypes[1]
|
||||
ms := protoimpl.X.MessageStateOf(protoimpl.Pointer(x))
|
||||
ms.StoreMessageInfo(mi)
|
||||
}
|
||||
|
||||
func (x *MultiUserServerConfig) String() string {
|
||||
return protoimpl.X.MessageStringOf(x)
|
||||
}
|
||||
|
||||
func (*MultiUserServerConfig) ProtoMessage() {}
|
||||
|
||||
func (x *MultiUserServerConfig) ProtoReflect() protoreflect.Message {
|
||||
mi := &file_proxy_shadowsocks_2022_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 MultiUserServerConfig.ProtoReflect.Descriptor instead.
|
||||
func (*MultiUserServerConfig) Descriptor() ([]byte, []int) {
|
||||
return file_proxy_shadowsocks_2022_config_proto_rawDescGZIP(), []int{1}
|
||||
}
|
||||
|
||||
func (x *MultiUserServerConfig) GetMethod() string {
|
||||
if x != nil {
|
||||
return x.Method
|
||||
}
|
||||
return ""
|
||||
}
|
||||
|
||||
func (x *MultiUserServerConfig) GetKey() string {
|
||||
if x != nil {
|
||||
return x.Key
|
||||
}
|
||||
return ""
|
||||
}
|
||||
|
||||
func (x *MultiUserServerConfig) GetUsers() []*protocol.User {
|
||||
if x != nil {
|
||||
return x.Users
|
||||
}
|
||||
return nil
|
||||
}
|
||||
|
||||
func (x *MultiUserServerConfig) GetNetwork() []net.Network {
|
||||
if x != nil {
|
||||
return x.Network
|
||||
}
|
||||
return nil
|
||||
}
|
||||
|
||||
type RelayDestination struct {
|
||||
state protoimpl.MessageState `protogen:"open.v1"`
|
||||
Key string `protobuf:"bytes,1,opt,name=key,proto3" json:"key,omitempty"`
|
||||
Address *net.IPOrDomain `protobuf:"bytes,2,opt,name=address,proto3" json:"address,omitempty"`
|
||||
Port uint32 `protobuf:"varint,3,opt,name=port,proto3" json:"port,omitempty"`
|
||||
Email string `protobuf:"bytes,4,opt,name=email,proto3" json:"email,omitempty"`
|
||||
Level int32 `protobuf:"varint,5,opt,name=level,proto3" json:"level,omitempty"`
|
||||
unknownFields protoimpl.UnknownFields
|
||||
sizeCache protoimpl.SizeCache
|
||||
}
|
||||
|
||||
func (x *RelayDestination) Reset() {
|
||||
*x = RelayDestination{}
|
||||
mi := &file_proxy_shadowsocks_2022_config_proto_msgTypes[2]
|
||||
ms := protoimpl.X.MessageStateOf(protoimpl.Pointer(x))
|
||||
ms.StoreMessageInfo(mi)
|
||||
}
|
||||
|
||||
func (x *RelayDestination) String() string {
|
||||
return protoimpl.X.MessageStringOf(x)
|
||||
}
|
||||
|
||||
func (*RelayDestination) ProtoMessage() {}
|
||||
|
||||
func (x *RelayDestination) ProtoReflect() protoreflect.Message {
|
||||
mi := &file_proxy_shadowsocks_2022_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 RelayDestination.ProtoReflect.Descriptor instead.
|
||||
func (*RelayDestination) Descriptor() ([]byte, []int) {
|
||||
return file_proxy_shadowsocks_2022_config_proto_rawDescGZIP(), []int{2}
|
||||
}
|
||||
|
||||
func (x *RelayDestination) GetKey() string {
|
||||
if x != nil {
|
||||
return x.Key
|
||||
}
|
||||
return ""
|
||||
}
|
||||
|
||||
func (x *RelayDestination) GetAddress() *net.IPOrDomain {
|
||||
if x != nil {
|
||||
return x.Address
|
||||
}
|
||||
return nil
|
||||
}
|
||||
|
||||
func (x *RelayDestination) GetPort() uint32 {
|
||||
if x != nil {
|
||||
return x.Port
|
||||
}
|
||||
return 0
|
||||
}
|
||||
|
||||
func (x *RelayDestination) GetEmail() string {
|
||||
if x != nil {
|
||||
return x.Email
|
||||
}
|
||||
return ""
|
||||
}
|
||||
|
||||
func (x *RelayDestination) GetLevel() int32 {
|
||||
if x != nil {
|
||||
return x.Level
|
||||
}
|
||||
return 0
|
||||
}
|
||||
|
||||
type RelayServerConfig struct {
|
||||
state protoimpl.MessageState `protogen:"open.v1"`
|
||||
Method string `protobuf:"bytes,1,opt,name=method,proto3" json:"method,omitempty"`
|
||||
Key string `protobuf:"bytes,2,opt,name=key,proto3" json:"key,omitempty"`
|
||||
Destinations []*RelayDestination `protobuf:"bytes,3,rep,name=destinations,proto3" json:"destinations,omitempty"`
|
||||
Network []net.Network `protobuf:"varint,4,rep,packed,name=network,proto3,enum=xray.common.net.Network" json:"network,omitempty"`
|
||||
unknownFields protoimpl.UnknownFields
|
||||
sizeCache protoimpl.SizeCache
|
||||
}
|
||||
|
||||
func (x *RelayServerConfig) Reset() {
|
||||
*x = RelayServerConfig{}
|
||||
mi := &file_proxy_shadowsocks_2022_config_proto_msgTypes[3]
|
||||
ms := protoimpl.X.MessageStateOf(protoimpl.Pointer(x))
|
||||
ms.StoreMessageInfo(mi)
|
||||
}
|
||||
|
||||
func (x *RelayServerConfig) String() string {
|
||||
return protoimpl.X.MessageStringOf(x)
|
||||
}
|
||||
|
||||
func (*RelayServerConfig) ProtoMessage() {}
|
||||
|
||||
func (x *RelayServerConfig) ProtoReflect() protoreflect.Message {
|
||||
mi := &file_proxy_shadowsocks_2022_config_proto_msgTypes[3]
|
||||
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 RelayServerConfig.ProtoReflect.Descriptor instead.
|
||||
func (*RelayServerConfig) Descriptor() ([]byte, []int) {
|
||||
return file_proxy_shadowsocks_2022_config_proto_rawDescGZIP(), []int{3}
|
||||
}
|
||||
|
||||
func (x *RelayServerConfig) GetMethod() string {
|
||||
if x != nil {
|
||||
return x.Method
|
||||
}
|
||||
return ""
|
||||
}
|
||||
|
||||
func (x *RelayServerConfig) GetKey() string {
|
||||
if x != nil {
|
||||
return x.Key
|
||||
}
|
||||
return ""
|
||||
}
|
||||
|
||||
func (x *RelayServerConfig) GetDestinations() []*RelayDestination {
|
||||
if x != nil {
|
||||
return x.Destinations
|
||||
}
|
||||
return nil
|
||||
}
|
||||
|
||||
func (x *RelayServerConfig) GetNetwork() []net.Network {
|
||||
if x != nil {
|
||||
return x.Network
|
||||
}
|
||||
return nil
|
||||
}
|
||||
|
||||
type Account struct {
|
||||
state protoimpl.MessageState `protogen:"open.v1"`
|
||||
Key string `protobuf:"bytes,1,opt,name=key,proto3" json:"key,omitempty"`
|
||||
unknownFields protoimpl.UnknownFields
|
||||
sizeCache protoimpl.SizeCache
|
||||
}
|
||||
|
||||
func (x *Account) Reset() {
|
||||
*x = Account{}
|
||||
mi := &file_proxy_shadowsocks_2022_config_proto_msgTypes[4]
|
||||
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_shadowsocks_2022_config_proto_msgTypes[4]
|
||||
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_shadowsocks_2022_config_proto_rawDescGZIP(), []int{4}
|
||||
}
|
||||
|
||||
func (x *Account) GetKey() string {
|
||||
if x != nil {
|
||||
return x.Key
|
||||
}
|
||||
return ""
|
||||
}
|
||||
|
||||
type ClientConfig struct {
|
||||
state protoimpl.MessageState `protogen:"open.v1"`
|
||||
Address *net.IPOrDomain `protobuf:"bytes,1,opt,name=address,proto3" json:"address,omitempty"`
|
||||
Port uint32 `protobuf:"varint,2,opt,name=port,proto3" json:"port,omitempty"`
|
||||
Method string `protobuf:"bytes,3,opt,name=method,proto3" json:"method,omitempty"`
|
||||
Key string `protobuf:"bytes,4,opt,name=key,proto3" json:"key,omitempty"`
|
||||
unknownFields protoimpl.UnknownFields
|
||||
sizeCache protoimpl.SizeCache
|
||||
}
|
||||
|
||||
func (x *ClientConfig) Reset() {
|
||||
*x = ClientConfig{}
|
||||
mi := &file_proxy_shadowsocks_2022_config_proto_msgTypes[5]
|
||||
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_shadowsocks_2022_config_proto_msgTypes[5]
|
||||
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_shadowsocks_2022_config_proto_rawDescGZIP(), []int{5}
|
||||
}
|
||||
|
||||
func (x *ClientConfig) GetAddress() *net.IPOrDomain {
|
||||
if x != nil {
|
||||
return x.Address
|
||||
}
|
||||
return nil
|
||||
}
|
||||
|
||||
func (x *ClientConfig) GetPort() uint32 {
|
||||
if x != nil {
|
||||
return x.Port
|
||||
}
|
||||
return 0
|
||||
}
|
||||
|
||||
func (x *ClientConfig) GetMethod() string {
|
||||
if x != nil {
|
||||
return x.Method
|
||||
}
|
||||
return ""
|
||||
}
|
||||
|
||||
func (x *ClientConfig) GetKey() string {
|
||||
if x != nil {
|
||||
return x.Key
|
||||
}
|
||||
return ""
|
||||
}
|
||||
|
||||
var File_proxy_shadowsocks_2022_config_proto protoreflect.FileDescriptor
|
||||
|
||||
const file_proxy_shadowsocks_2022_config_proto_rawDesc = "" +
|
||||
"\n" +
|
||||
"#proxy/shadowsocks_2022/config.proto\x12\x1bxray.proxy.shadowsocks_2022\x1a\x18common/net/network.proto\x1a\x18common/net/address.proto\x1a\x1acommon/protocol/user.proto\"\x98\x01\n" +
|
||||
"\fServerConfig\x12\x16\n" +
|
||||
"\x06method\x18\x01 \x01(\tR\x06method\x12\x10\n" +
|
||||
"\x03key\x18\x02 \x01(\tR\x03key\x12\x14\n" +
|
||||
"\x05email\x18\x03 \x01(\tR\x05email\x12\x14\n" +
|
||||
"\x05level\x18\x04 \x01(\x05R\x05level\x122\n" +
|
||||
"\anetwork\x18\x05 \x03(\x0e2\x18.xray.common.net.NetworkR\anetwork\"\xa7\x01\n" +
|
||||
"\x15MultiUserServerConfig\x12\x16\n" +
|
||||
"\x06method\x18\x01 \x01(\tR\x06method\x12\x10\n" +
|
||||
"\x03key\x18\x02 \x01(\tR\x03key\x120\n" +
|
||||
"\x05users\x18\x03 \x03(\v2\x1a.xray.common.protocol.UserR\x05users\x122\n" +
|
||||
"\anetwork\x18\x04 \x03(\x0e2\x18.xray.common.net.NetworkR\anetwork\"\x9b\x01\n" +
|
||||
"\x10RelayDestination\x12\x10\n" +
|
||||
"\x03key\x18\x01 \x01(\tR\x03key\x125\n" +
|
||||
"\aaddress\x18\x02 \x01(\v2\x1b.xray.common.net.IPOrDomainR\aaddress\x12\x12\n" +
|
||||
"\x04port\x18\x03 \x01(\rR\x04port\x12\x14\n" +
|
||||
"\x05email\x18\x04 \x01(\tR\x05email\x12\x14\n" +
|
||||
"\x05level\x18\x05 \x01(\x05R\x05level\"\xc4\x01\n" +
|
||||
"\x11RelayServerConfig\x12\x16\n" +
|
||||
"\x06method\x18\x01 \x01(\tR\x06method\x12\x10\n" +
|
||||
"\x03key\x18\x02 \x01(\tR\x03key\x12Q\n" +
|
||||
"\fdestinations\x18\x03 \x03(\v2-.xray.proxy.shadowsocks_2022.RelayDestinationR\fdestinations\x122\n" +
|
||||
"\anetwork\x18\x04 \x03(\x0e2\x18.xray.common.net.NetworkR\anetwork\"\x1b\n" +
|
||||
"\aAccount\x12\x10\n" +
|
||||
"\x03key\x18\x01 \x01(\tR\x03key\"\x83\x01\n" +
|
||||
"\fClientConfig\x125\n" +
|
||||
"\aaddress\x18\x01 \x01(\v2\x1b.xray.common.net.IPOrDomainR\aaddress\x12\x12\n" +
|
||||
"\x04port\x18\x02 \x01(\rR\x04port\x12\x16\n" +
|
||||
"\x06method\x18\x03 \x01(\tR\x06method\x12\x10\n" +
|
||||
"\x03key\x18\x04 \x01(\tR\x03keyBr\n" +
|
||||
"\x1fcom.xray.proxy.shadowsocks_2022P\x01Z0github.com/xtls/xray-core/proxy/shadowsocks_2022\xaa\x02\x1aXray.Proxy.Shadowsocks2022b\x06proto3"
|
||||
|
||||
var (
|
||||
file_proxy_shadowsocks_2022_config_proto_rawDescOnce sync.Once
|
||||
file_proxy_shadowsocks_2022_config_proto_rawDescData []byte
|
||||
)
|
||||
|
||||
func file_proxy_shadowsocks_2022_config_proto_rawDescGZIP() []byte {
|
||||
file_proxy_shadowsocks_2022_config_proto_rawDescOnce.Do(func() {
|
||||
file_proxy_shadowsocks_2022_config_proto_rawDescData = protoimpl.X.CompressGZIP(unsafe.Slice(unsafe.StringData(file_proxy_shadowsocks_2022_config_proto_rawDesc), len(file_proxy_shadowsocks_2022_config_proto_rawDesc)))
|
||||
})
|
||||
return file_proxy_shadowsocks_2022_config_proto_rawDescData
|
||||
}
|
||||
|
||||
var file_proxy_shadowsocks_2022_config_proto_msgTypes = make([]protoimpl.MessageInfo, 6)
|
||||
var file_proxy_shadowsocks_2022_config_proto_goTypes = []any{
|
||||
(*ServerConfig)(nil), // 0: xray.proxy.shadowsocks_2022.ServerConfig
|
||||
(*MultiUserServerConfig)(nil), // 1: xray.proxy.shadowsocks_2022.MultiUserServerConfig
|
||||
(*RelayDestination)(nil), // 2: xray.proxy.shadowsocks_2022.RelayDestination
|
||||
(*RelayServerConfig)(nil), // 3: xray.proxy.shadowsocks_2022.RelayServerConfig
|
||||
(*Account)(nil), // 4: xray.proxy.shadowsocks_2022.Account
|
||||
(*ClientConfig)(nil), // 5: xray.proxy.shadowsocks_2022.ClientConfig
|
||||
(net.Network)(0), // 6: xray.common.net.Network
|
||||
(*protocol.User)(nil), // 7: xray.common.protocol.User
|
||||
(*net.IPOrDomain)(nil), // 8: xray.common.net.IPOrDomain
|
||||
}
|
||||
var file_proxy_shadowsocks_2022_config_proto_depIdxs = []int32{
|
||||
6, // 0: xray.proxy.shadowsocks_2022.ServerConfig.network:type_name -> xray.common.net.Network
|
||||
7, // 1: xray.proxy.shadowsocks_2022.MultiUserServerConfig.users:type_name -> xray.common.protocol.User
|
||||
6, // 2: xray.proxy.shadowsocks_2022.MultiUserServerConfig.network:type_name -> xray.common.net.Network
|
||||
8, // 3: xray.proxy.shadowsocks_2022.RelayDestination.address:type_name -> xray.common.net.IPOrDomain
|
||||
2, // 4: xray.proxy.shadowsocks_2022.RelayServerConfig.destinations:type_name -> xray.proxy.shadowsocks_2022.RelayDestination
|
||||
6, // 5: xray.proxy.shadowsocks_2022.RelayServerConfig.network:type_name -> xray.common.net.Network
|
||||
8, // 6: xray.proxy.shadowsocks_2022.ClientConfig.address:type_name -> xray.common.net.IPOrDomain
|
||||
7, // [7:7] is the sub-list for method output_type
|
||||
7, // [7:7] is the sub-list for method input_type
|
||||
7, // [7:7] is the sub-list for extension type_name
|
||||
7, // [7:7] is the sub-list for extension extendee
|
||||
0, // [0:7] is the sub-list for field type_name
|
||||
}
|
||||
|
||||
func init() { file_proxy_shadowsocks_2022_config_proto_init() }
|
||||
func file_proxy_shadowsocks_2022_config_proto_init() {
|
||||
if File_proxy_shadowsocks_2022_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_shadowsocks_2022_config_proto_rawDesc), len(file_proxy_shadowsocks_2022_config_proto_rawDesc)),
|
||||
NumEnums: 0,
|
||||
NumMessages: 6,
|
||||
NumExtensions: 0,
|
||||
NumServices: 0,
|
||||
},
|
||||
GoTypes: file_proxy_shadowsocks_2022_config_proto_goTypes,
|
||||
DependencyIndexes: file_proxy_shadowsocks_2022_config_proto_depIdxs,
|
||||
MessageInfos: file_proxy_shadowsocks_2022_config_proto_msgTypes,
|
||||
}.Build()
|
||||
File_proxy_shadowsocks_2022_config_proto = out.File
|
||||
file_proxy_shadowsocks_2022_config_proto_goTypes = nil
|
||||
file_proxy_shadowsocks_2022_config_proto_depIdxs = nil
|
||||
}
|
||||
@@ -0,0 +1,52 @@
|
||||
syntax = "proto3";
|
||||
|
||||
package xray.proxy.shadowsocks_2022;
|
||||
option csharp_namespace = "Xray.Proxy.Shadowsocks2022";
|
||||
option go_package = "github.com/xtls/xray-core/proxy/shadowsocks_2022";
|
||||
option java_package = "com.xray.proxy.shadowsocks_2022";
|
||||
option java_multiple_files = true;
|
||||
|
||||
import "common/net/network.proto";
|
||||
import "common/net/address.proto";
|
||||
import "common/protocol/user.proto";
|
||||
|
||||
message ServerConfig {
|
||||
string method = 1;
|
||||
string key = 2;
|
||||
string email = 3;
|
||||
int32 level = 4;
|
||||
repeated xray.common.net.Network network = 5;
|
||||
}
|
||||
|
||||
message MultiUserServerConfig {
|
||||
string method = 1;
|
||||
string key = 2;
|
||||
repeated xray.common.protocol.User users = 3;
|
||||
repeated xray.common.net.Network network = 4;
|
||||
}
|
||||
|
||||
message RelayDestination {
|
||||
string key = 1;
|
||||
xray.common.net.IPOrDomain address = 2;
|
||||
uint32 port = 3;
|
||||
string email = 4;
|
||||
int32 level = 5;
|
||||
}
|
||||
|
||||
message RelayServerConfig {
|
||||
string method = 1;
|
||||
string key = 2;
|
||||
repeated RelayDestination destinations = 3;
|
||||
repeated xray.common.net.Network network = 4;
|
||||
}
|
||||
|
||||
message Account {
|
||||
string key = 1;
|
||||
}
|
||||
|
||||
message ClientConfig {
|
||||
xray.common.net.IPOrDomain address = 1;
|
||||
uint32 port = 2;
|
||||
string method = 3;
|
||||
string key = 4;
|
||||
}
|
||||
@@ -0,0 +1,183 @@
|
||||
package shadowsocks_2022
|
||||
|
||||
import (
|
||||
"context"
|
||||
"time"
|
||||
|
||||
shadowsocks "github.com/sagernet/sing-shadowsocks"
|
||||
"github.com/sagernet/sing-shadowsocks/shadowaead_2022"
|
||||
C "github.com/sagernet/sing/common"
|
||||
B "github.com/sagernet/sing/common/buf"
|
||||
"github.com/sagernet/sing/common/bufio"
|
||||
E "github.com/sagernet/sing/common/exceptions"
|
||||
M "github.com/sagernet/sing/common/metadata"
|
||||
N "github.com/sagernet/sing/common/network"
|
||||
"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/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/common/signal"
|
||||
"github.com/xtls/xray-core/common/singbridge"
|
||||
"github.com/xtls/xray-core/features/routing"
|
||||
"github.com/xtls/xray-core/transport/internet/stat"
|
||||
)
|
||||
|
||||
func init() {
|
||||
common.Must(common.RegisterConfig((*ServerConfig)(nil), func(ctx context.Context, config interface{}) (interface{}, error) {
|
||||
return NewServer(ctx, config.(*ServerConfig))
|
||||
}))
|
||||
}
|
||||
|
||||
type Inbound struct {
|
||||
networks []net.Network
|
||||
service shadowsocks.Service
|
||||
email string
|
||||
level int
|
||||
}
|
||||
|
||||
func NewServer(ctx context.Context, config *ServerConfig) (*Inbound, error) {
|
||||
networks := config.Network
|
||||
if len(networks) == 0 {
|
||||
networks = []net.Network{
|
||||
net.Network_TCP,
|
||||
net.Network_UDP,
|
||||
}
|
||||
}
|
||||
inbound := &Inbound{
|
||||
networks: networks,
|
||||
email: config.Email,
|
||||
level: int(config.Level),
|
||||
}
|
||||
if !C.Contains(shadowaead_2022.List, config.Method) {
|
||||
return nil, errors.New("unsupported method ", config.Method)
|
||||
}
|
||||
service, err := shadowaead_2022.NewServiceWithPassword(config.Method, config.Key, 500, inbound, nil)
|
||||
if err != nil {
|
||||
return nil, errors.New("create service").Base(err)
|
||||
}
|
||||
inbound.service = service
|
||||
return inbound, nil
|
||||
}
|
||||
|
||||
func (i *Inbound) Network() []net.Network {
|
||||
return i.networks
|
||||
}
|
||||
|
||||
func (i *Inbound) Process(ctx context.Context, network net.Network, connection stat.Connection, dispatcher routing.Dispatcher) error {
|
||||
inbound := session.InboundFromContext(ctx)
|
||||
inbound.Name = "shadowsocks-2022"
|
||||
inbound.CanSpliceCopy = 3
|
||||
|
||||
var metadata M.Metadata
|
||||
if inbound.Source.IsValid() {
|
||||
metadata.Source = M.ParseSocksaddr(inbound.Source.NetAddr())
|
||||
}
|
||||
|
||||
ctx = session.ContextWithDispatcher(ctx, dispatcher)
|
||||
|
||||
if network == net.Network_TCP {
|
||||
return singbridge.ReturnError(i.service.NewConnection(ctx, connection, metadata))
|
||||
} else {
|
||||
reader := buf.NewReader(connection)
|
||||
pc := &natPacketConn{connection}
|
||||
for {
|
||||
mb, err := reader.ReadMultiBuffer()
|
||||
if err != nil {
|
||||
buf.ReleaseMulti(mb)
|
||||
return singbridge.ReturnError(err)
|
||||
}
|
||||
for _, buffer := range mb {
|
||||
packet := B.As(buffer.Bytes()).ToOwned()
|
||||
buffer.Release()
|
||||
err = i.service.NewPacket(ctx, pc, packet, metadata)
|
||||
if err != nil {
|
||||
packet.Release()
|
||||
buf.ReleaseMulti(mb)
|
||||
return err
|
||||
}
|
||||
}
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
func (i *Inbound) NewConnection(ctx context.Context, conn net.Conn, metadata M.Metadata) error {
|
||||
inbound := session.InboundFromContext(ctx)
|
||||
inbound.User = &protocol.MemoryUser{
|
||||
Email: i.email,
|
||||
Level: uint32(i.level),
|
||||
}
|
||||
ctx = log.ContextWithAccessMessage(ctx, &log.AccessMessage{
|
||||
From: metadata.Source,
|
||||
To: metadata.Destination,
|
||||
Status: log.AccessAccepted,
|
||||
Email: i.email,
|
||||
})
|
||||
errors.LogInfo(ctx, "tunnelling request to tcp:", metadata.Destination)
|
||||
dispatcher := session.DispatcherFromContext(ctx)
|
||||
destination, err := singbridge.ToDestination(metadata.Destination, net.Network_TCP)
|
||||
if err != nil {
|
||||
return err
|
||||
}
|
||||
link, err := dispatcher.Dispatch(ctx, destination)
|
||||
if err != nil {
|
||||
return err
|
||||
}
|
||||
return singbridge.CopyConn(ctx, nil, link, conn)
|
||||
}
|
||||
|
||||
func (i *Inbound) NewPacketConnection(ctx context.Context, conn N.PacketConn, metadata M.Metadata) error {
|
||||
inbound := session.InboundFromContext(ctx)
|
||||
inbound.User = &protocol.MemoryUser{
|
||||
Email: i.email,
|
||||
Level: uint32(i.level),
|
||||
}
|
||||
ctx = log.ContextWithAccessMessage(ctx, &log.AccessMessage{
|
||||
From: metadata.Source,
|
||||
To: metadata.Destination,
|
||||
Status: log.AccessAccepted,
|
||||
Email: i.email,
|
||||
})
|
||||
errors.LogInfo(ctx, "tunnelling request to udp:", metadata.Destination)
|
||||
dispatcher := session.DispatcherFromContext(ctx)
|
||||
destination, err := singbridge.ToDestination(metadata.Destination, net.Network_UDP)
|
||||
if err != nil {
|
||||
return err
|
||||
}
|
||||
link, err := dispatcher.Dispatch(ctx, destination)
|
||||
if err != nil {
|
||||
return err
|
||||
}
|
||||
outConn := &singbridge.PacketConnWrapper{
|
||||
Reader: link.Reader,
|
||||
Writer: link.Writer,
|
||||
Dest: destination,
|
||||
T: signal.CancelAfterInactivity(ctx, func() {
|
||||
common.Interrupt(link.Reader)
|
||||
}, 300*time.Second),
|
||||
}
|
||||
return bufio.CopyPacketConn(ctx, conn, outConn)
|
||||
}
|
||||
|
||||
func (i *Inbound) NewError(ctx context.Context, err error) {
|
||||
if E.IsClosed(err) {
|
||||
return
|
||||
}
|
||||
errors.LogWarning(ctx, err.Error())
|
||||
}
|
||||
|
||||
type natPacketConn struct {
|
||||
net.Conn
|
||||
}
|
||||
|
||||
func (c *natPacketConn) ReadPacket(buffer *B.Buffer) (addr M.Socksaddr, err error) {
|
||||
_, err = buffer.ReadFrom(c)
|
||||
return
|
||||
}
|
||||
|
||||
func (c *natPacketConn) WritePacket(buffer *B.Buffer, addr M.Socksaddr) error {
|
||||
_, err := buffer.WriteTo(c)
|
||||
return err
|
||||
}
|
||||
@@ -0,0 +1,290 @@
|
||||
package shadowsocks_2022
|
||||
|
||||
import (
|
||||
"context"
|
||||
"encoding/base64"
|
||||
"strconv"
|
||||
"strings"
|
||||
"sync"
|
||||
"time"
|
||||
|
||||
"github.com/sagernet/sing-shadowsocks/shadowaead_2022"
|
||||
C "github.com/sagernet/sing/common"
|
||||
A "github.com/sagernet/sing/common/auth"
|
||||
B "github.com/sagernet/sing/common/buf"
|
||||
"github.com/sagernet/sing/common/bufio"
|
||||
E "github.com/sagernet/sing/common/exceptions"
|
||||
M "github.com/sagernet/sing/common/metadata"
|
||||
N "github.com/sagernet/sing/common/network"
|
||||
"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/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/common/signal"
|
||||
"github.com/xtls/xray-core/common/singbridge"
|
||||
"github.com/xtls/xray-core/common/uuid"
|
||||
"github.com/xtls/xray-core/features/routing"
|
||||
"github.com/xtls/xray-core/transport/internet/stat"
|
||||
)
|
||||
|
||||
func init() {
|
||||
common.Must(common.RegisterConfig((*MultiUserServerConfig)(nil), func(ctx context.Context, config interface{}) (interface{}, error) {
|
||||
return NewMultiServer(ctx, config.(*MultiUserServerConfig))
|
||||
}))
|
||||
}
|
||||
|
||||
type MultiUserInbound struct {
|
||||
sync.Mutex
|
||||
networks []net.Network
|
||||
users []*protocol.MemoryUser
|
||||
service *shadowaead_2022.MultiService[int]
|
||||
}
|
||||
|
||||
func NewMultiServer(ctx context.Context, config *MultiUserServerConfig) (*MultiUserInbound, error) {
|
||||
networks := config.Network
|
||||
if len(networks) == 0 {
|
||||
networks = []net.Network{
|
||||
net.Network_TCP,
|
||||
net.Network_UDP,
|
||||
}
|
||||
}
|
||||
memUsers := []*protocol.MemoryUser{}
|
||||
for i, user := range config.Users {
|
||||
if user.Email == "" {
|
||||
u := uuid.New()
|
||||
user.Email = "unnamed-user-" + strconv.Itoa(i) + "-" + u.String()
|
||||
}
|
||||
u, err := user.ToMemoryUser()
|
||||
if err != nil {
|
||||
return nil, errors.New("failed to get shadowsocks user").Base(err).AtError()
|
||||
}
|
||||
memUsers = append(memUsers, u)
|
||||
}
|
||||
|
||||
inbound := &MultiUserInbound{
|
||||
networks: networks,
|
||||
users: memUsers,
|
||||
}
|
||||
if config.Key == "" {
|
||||
return nil, errors.New("missing key")
|
||||
}
|
||||
psk, err := base64.StdEncoding.DecodeString(config.Key)
|
||||
if err != nil {
|
||||
return nil, errors.New("parse config").Base(err)
|
||||
}
|
||||
service, err := shadowaead_2022.NewMultiService[int](config.Method, psk, 500, inbound, nil)
|
||||
if err != nil {
|
||||
return nil, errors.New("create service").Base(err)
|
||||
}
|
||||
err = service.UpdateUsersWithPasswords(
|
||||
C.MapIndexed(memUsers, func(index int, it *protocol.MemoryUser) int { return index }),
|
||||
C.Map(memUsers, func(it *protocol.MemoryUser) string { return it.Account.(*MemoryAccount).Key }),
|
||||
)
|
||||
if err != nil {
|
||||
return nil, errors.New("create service").Base(err)
|
||||
}
|
||||
|
||||
inbound.service = service
|
||||
return inbound, nil
|
||||
}
|
||||
|
||||
// AddUser implements proxy.UserManager.AddUser().
|
||||
func (i *MultiUserInbound) AddUser(ctx context.Context, u *protocol.MemoryUser) error {
|
||||
i.Lock()
|
||||
defer i.Unlock()
|
||||
|
||||
if u.Email != "" {
|
||||
for idx := range i.users {
|
||||
if i.users[idx].Email == u.Email {
|
||||
return errors.New("User ", u.Email, " already exists.")
|
||||
}
|
||||
}
|
||||
}
|
||||
i.users = append(i.users, u)
|
||||
|
||||
// sync to multi service
|
||||
// Considering implements shadowsocks2022 in xray-core may have better performance.
|
||||
i.service.UpdateUsersWithPasswords(
|
||||
C.MapIndexed(i.users, func(index int, it *protocol.MemoryUser) int { return index }),
|
||||
C.Map(i.users, func(it *protocol.MemoryUser) string { return it.Account.(*MemoryAccount).Key }),
|
||||
)
|
||||
|
||||
return nil
|
||||
}
|
||||
|
||||
// RemoveUser implements proxy.UserManager.RemoveUser().
|
||||
func (i *MultiUserInbound) RemoveUser(ctx context.Context, email string) error {
|
||||
if email == "" {
|
||||
return errors.New("Email must not be empty.")
|
||||
}
|
||||
|
||||
i.Lock()
|
||||
defer i.Unlock()
|
||||
|
||||
idx := -1
|
||||
for ii, u := range i.users {
|
||||
if strings.EqualFold(u.Email, email) {
|
||||
idx = ii
|
||||
break
|
||||
}
|
||||
}
|
||||
|
||||
if idx == -1 {
|
||||
return errors.New("User ", email, " not found.")
|
||||
}
|
||||
|
||||
ulen := len(i.users)
|
||||
|
||||
i.users[idx] = i.users[ulen-1]
|
||||
i.users[ulen-1] = nil
|
||||
i.users = i.users[:ulen-1]
|
||||
|
||||
// sync to multi service
|
||||
// Considering implements shadowsocks2022 in xray-core may have better performance.
|
||||
i.service.UpdateUsersWithPasswords(
|
||||
C.MapIndexed(i.users, func(index int, it *protocol.MemoryUser) int { return index }),
|
||||
C.Map(i.users, func(it *protocol.MemoryUser) string { return it.Account.(*MemoryAccount).Key }),
|
||||
)
|
||||
|
||||
return nil
|
||||
}
|
||||
|
||||
// GetUser implements proxy.UserManager.GetUser().
|
||||
func (i *MultiUserInbound) GetUser(ctx context.Context, email string) *protocol.MemoryUser {
|
||||
if email == "" {
|
||||
return nil
|
||||
}
|
||||
|
||||
i.Lock()
|
||||
defer i.Unlock()
|
||||
|
||||
for _, u := range i.users {
|
||||
if strings.EqualFold(u.Email, email) {
|
||||
return u
|
||||
}
|
||||
}
|
||||
return nil
|
||||
}
|
||||
|
||||
// GetUsers implements proxy.UserManager.GetUsers().
|
||||
func (i *MultiUserInbound) GetUsers(ctx context.Context) []*protocol.MemoryUser {
|
||||
i.Lock()
|
||||
defer i.Unlock()
|
||||
dst := make([]*protocol.MemoryUser, len(i.users))
|
||||
copy(dst, i.users)
|
||||
return dst
|
||||
}
|
||||
|
||||
// GetUsersCount implements proxy.UserManager.GetUsersCount().
|
||||
func (i *MultiUserInbound) GetUsersCount(context.Context) int64 {
|
||||
i.Lock()
|
||||
defer i.Unlock()
|
||||
return int64(len(i.users))
|
||||
}
|
||||
|
||||
func (i *MultiUserInbound) Network() []net.Network {
|
||||
return i.networks
|
||||
}
|
||||
|
||||
func (i *MultiUserInbound) Process(ctx context.Context, network net.Network, connection stat.Connection, dispatcher routing.Dispatcher) error {
|
||||
inbound := session.InboundFromContext(ctx)
|
||||
inbound.Name = "shadowsocks-2022-multi"
|
||||
inbound.CanSpliceCopy = 3
|
||||
|
||||
var metadata M.Metadata
|
||||
if inbound.Source.IsValid() {
|
||||
metadata.Source = M.ParseSocksaddr(inbound.Source.NetAddr())
|
||||
}
|
||||
|
||||
ctx = session.ContextWithDispatcher(ctx, dispatcher)
|
||||
|
||||
if network == net.Network_TCP {
|
||||
return singbridge.ReturnError(i.service.NewConnection(ctx, connection, metadata))
|
||||
} else {
|
||||
reader := buf.NewReader(connection)
|
||||
pc := &natPacketConn{connection}
|
||||
for {
|
||||
mb, err := reader.ReadMultiBuffer()
|
||||
if err != nil {
|
||||
buf.ReleaseMulti(mb)
|
||||
return singbridge.ReturnError(err)
|
||||
}
|
||||
for _, buffer := range mb {
|
||||
packet := B.As(buffer.Bytes()).ToOwned()
|
||||
buffer.Release()
|
||||
err = i.service.NewPacket(ctx, pc, packet, metadata)
|
||||
if err != nil {
|
||||
packet.Release()
|
||||
buf.ReleaseMulti(mb)
|
||||
return err
|
||||
}
|
||||
}
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
func (i *MultiUserInbound) NewConnection(ctx context.Context, conn net.Conn, metadata M.Metadata) error {
|
||||
inbound := session.InboundFromContext(ctx)
|
||||
userInt, _ := A.UserFromContext[int](ctx)
|
||||
user := i.users[userInt]
|
||||
inbound.User = user
|
||||
ctx = log.ContextWithAccessMessage(ctx, &log.AccessMessage{
|
||||
From: metadata.Source,
|
||||
To: metadata.Destination,
|
||||
Status: log.AccessAccepted,
|
||||
Email: user.Email,
|
||||
})
|
||||
errors.LogInfo(ctx, "tunnelling request to tcp:", metadata.Destination)
|
||||
dispatcher := session.DispatcherFromContext(ctx)
|
||||
destination, err := singbridge.ToDestination(metadata.Destination, net.Network_TCP)
|
||||
if err != nil {
|
||||
return err
|
||||
}
|
||||
link, err := dispatcher.Dispatch(ctx, destination)
|
||||
if err != nil {
|
||||
return err
|
||||
}
|
||||
return singbridge.CopyConn(ctx, conn, link, conn)
|
||||
}
|
||||
|
||||
func (i *MultiUserInbound) NewPacketConnection(ctx context.Context, conn N.PacketConn, metadata M.Metadata) error {
|
||||
inbound := session.InboundFromContext(ctx)
|
||||
userInt, _ := A.UserFromContext[int](ctx)
|
||||
user := i.users[userInt]
|
||||
inbound.User = user
|
||||
ctx = log.ContextWithAccessMessage(ctx, &log.AccessMessage{
|
||||
From: metadata.Source,
|
||||
To: metadata.Destination,
|
||||
Status: log.AccessAccepted,
|
||||
Email: user.Email,
|
||||
})
|
||||
errors.LogInfo(ctx, "tunnelling request to udp:", metadata.Destination)
|
||||
dispatcher := session.DispatcherFromContext(ctx)
|
||||
destination, err := singbridge.ToDestination(metadata.Destination, net.Network_UDP)
|
||||
if err != nil {
|
||||
return err
|
||||
}
|
||||
link, err := dispatcher.Dispatch(ctx, destination)
|
||||
if err != nil {
|
||||
return err
|
||||
}
|
||||
outConn := &singbridge.PacketConnWrapper{
|
||||
Reader: link.Reader,
|
||||
Writer: link.Writer,
|
||||
Dest: destination,
|
||||
T: signal.CancelAfterInactivity(ctx, func() {
|
||||
common.Interrupt(link.Reader)
|
||||
}, 300*time.Second),
|
||||
}
|
||||
return bufio.CopyPacketConn(ctx, conn, outConn)
|
||||
}
|
||||
|
||||
func (i *MultiUserInbound) NewError(ctx context.Context, err error) {
|
||||
if E.IsClosed(err) {
|
||||
return
|
||||
}
|
||||
errors.LogWarning(ctx, err.Error())
|
||||
}
|
||||
@@ -0,0 +1,194 @@
|
||||
package shadowsocks_2022
|
||||
|
||||
import (
|
||||
"context"
|
||||
"strconv"
|
||||
"strings"
|
||||
"time"
|
||||
|
||||
"github.com/sagernet/sing-shadowsocks/shadowaead_2022"
|
||||
C "github.com/sagernet/sing/common"
|
||||
A "github.com/sagernet/sing/common/auth"
|
||||
B "github.com/sagernet/sing/common/buf"
|
||||
"github.com/sagernet/sing/common/bufio"
|
||||
E "github.com/sagernet/sing/common/exceptions"
|
||||
M "github.com/sagernet/sing/common/metadata"
|
||||
N "github.com/sagernet/sing/common/network"
|
||||
"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/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/common/signal"
|
||||
"github.com/xtls/xray-core/common/singbridge"
|
||||
"github.com/xtls/xray-core/common/uuid"
|
||||
"github.com/xtls/xray-core/features/routing"
|
||||
"github.com/xtls/xray-core/transport/internet/stat"
|
||||
)
|
||||
|
||||
func init() {
|
||||
common.Must(common.RegisterConfig((*RelayServerConfig)(nil), func(ctx context.Context, config interface{}) (interface{}, error) {
|
||||
return NewRelayServer(ctx, config.(*RelayServerConfig))
|
||||
}))
|
||||
}
|
||||
|
||||
type RelayInbound struct {
|
||||
networks []net.Network
|
||||
destinations []*RelayDestination
|
||||
service *shadowaead_2022.RelayService[int]
|
||||
}
|
||||
|
||||
func NewRelayServer(ctx context.Context, config *RelayServerConfig) (*RelayInbound, error) {
|
||||
networks := config.Network
|
||||
if len(networks) == 0 {
|
||||
networks = []net.Network{
|
||||
net.Network_TCP,
|
||||
net.Network_UDP,
|
||||
}
|
||||
}
|
||||
inbound := &RelayInbound{
|
||||
networks: networks,
|
||||
destinations: config.Destinations,
|
||||
}
|
||||
if !C.Contains(shadowaead_2022.List, config.Method) || !strings.Contains(config.Method, "aes") {
|
||||
return nil, errors.New("unsupported method ", config.Method)
|
||||
}
|
||||
service, err := shadowaead_2022.NewRelayServiceWithPassword[int](config.Method, config.Key, 500, inbound)
|
||||
if err != nil {
|
||||
return nil, errors.New("create service").Base(err)
|
||||
}
|
||||
|
||||
for i, destination := range config.Destinations {
|
||||
if destination.Email == "" {
|
||||
u := uuid.New()
|
||||
destination.Email = "unnamed-destination-" + strconv.Itoa(i) + "-" + u.String()
|
||||
}
|
||||
}
|
||||
err = service.UpdateUsersWithPasswords(
|
||||
C.MapIndexed(config.Destinations, func(index int, it *RelayDestination) int { return index }),
|
||||
C.Map(config.Destinations, func(it *RelayDestination) string { return it.Key }),
|
||||
C.Map(config.Destinations, func(it *RelayDestination) M.Socksaddr {
|
||||
return singbridge.ToSocksaddr(net.Destination{
|
||||
Address: it.Address.AsAddress(),
|
||||
Port: net.Port(it.Port),
|
||||
})
|
||||
}),
|
||||
)
|
||||
if err != nil {
|
||||
return nil, errors.New("create service").Base(err)
|
||||
}
|
||||
inbound.service = service
|
||||
return inbound, nil
|
||||
}
|
||||
|
||||
func (i *RelayInbound) Network() []net.Network {
|
||||
return i.networks
|
||||
}
|
||||
|
||||
func (i *RelayInbound) Process(ctx context.Context, network net.Network, connection stat.Connection, dispatcher routing.Dispatcher) error {
|
||||
inbound := session.InboundFromContext(ctx)
|
||||
inbound.Name = "shadowsocks-2022-relay"
|
||||
inbound.CanSpliceCopy = 3
|
||||
|
||||
var metadata M.Metadata
|
||||
if inbound.Source.IsValid() {
|
||||
metadata.Source = M.ParseSocksaddr(inbound.Source.NetAddr())
|
||||
}
|
||||
|
||||
ctx = session.ContextWithDispatcher(ctx, dispatcher)
|
||||
|
||||
if network == net.Network_TCP {
|
||||
return singbridge.ReturnError(i.service.NewConnection(ctx, connection, metadata))
|
||||
} else {
|
||||
reader := buf.NewReader(connection)
|
||||
pc := &natPacketConn{connection}
|
||||
for {
|
||||
mb, err := reader.ReadMultiBuffer()
|
||||
if err != nil {
|
||||
buf.ReleaseMulti(mb)
|
||||
return singbridge.ReturnError(err)
|
||||
}
|
||||
for _, buffer := range mb {
|
||||
packet := B.As(buffer.Bytes()).ToOwned()
|
||||
buffer.Release()
|
||||
err = i.service.NewPacket(ctx, pc, packet, metadata)
|
||||
if err != nil {
|
||||
packet.Release()
|
||||
buf.ReleaseMulti(mb)
|
||||
return err
|
||||
}
|
||||
}
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
func (i *RelayInbound) NewConnection(ctx context.Context, conn net.Conn, metadata M.Metadata) error {
|
||||
inbound := session.InboundFromContext(ctx)
|
||||
userInt, _ := A.UserFromContext[int](ctx)
|
||||
user := i.destinations[userInt]
|
||||
inbound.User = &protocol.MemoryUser{
|
||||
Email: user.Email,
|
||||
Level: uint32(user.Level),
|
||||
}
|
||||
ctx = log.ContextWithAccessMessage(ctx, &log.AccessMessage{
|
||||
From: metadata.Source,
|
||||
To: metadata.Destination,
|
||||
Status: log.AccessAccepted,
|
||||
Email: user.Email,
|
||||
})
|
||||
errors.LogInfo(ctx, "tunnelling request to tcp:", metadata.Destination)
|
||||
dispatcher := session.DispatcherFromContext(ctx)
|
||||
destination, err := singbridge.ToDestination(metadata.Destination, net.Network_TCP)
|
||||
if err != nil {
|
||||
return err
|
||||
}
|
||||
link, err := dispatcher.Dispatch(ctx, destination)
|
||||
if err != nil {
|
||||
return err
|
||||
}
|
||||
return singbridge.CopyConn(ctx, nil, link, conn)
|
||||
}
|
||||
|
||||
func (i *RelayInbound) NewPacketConnection(ctx context.Context, conn N.PacketConn, metadata M.Metadata) error {
|
||||
inbound := session.InboundFromContext(ctx)
|
||||
userInt, _ := A.UserFromContext[int](ctx)
|
||||
user := i.destinations[userInt]
|
||||
inbound.User = &protocol.MemoryUser{
|
||||
Email: user.Email,
|
||||
Level: uint32(user.Level),
|
||||
}
|
||||
ctx = log.ContextWithAccessMessage(ctx, &log.AccessMessage{
|
||||
From: metadata.Source,
|
||||
To: metadata.Destination,
|
||||
Status: log.AccessAccepted,
|
||||
Email: user.Email,
|
||||
})
|
||||
errors.LogInfo(ctx, "tunnelling request to udp:", metadata.Destination)
|
||||
dispatcher := session.DispatcherFromContext(ctx)
|
||||
destination, err := singbridge.ToDestination(metadata.Destination, net.Network_UDP)
|
||||
if err != nil {
|
||||
return err
|
||||
}
|
||||
link, err := dispatcher.Dispatch(ctx, destination)
|
||||
if err != nil {
|
||||
return err
|
||||
}
|
||||
outConn := &singbridge.PacketConnWrapper{
|
||||
Reader: link.Reader,
|
||||
Writer: link.Writer,
|
||||
Dest: destination,
|
||||
T: signal.CancelAfterInactivity(ctx, func() {
|
||||
common.Interrupt(link.Reader)
|
||||
}, 300*time.Second),
|
||||
}
|
||||
return bufio.CopyPacketConn(ctx, conn, outConn)
|
||||
}
|
||||
|
||||
func (i *RelayInbound) NewError(ctx context.Context, err error) {
|
||||
if E.IsClosed(err) {
|
||||
return
|
||||
}
|
||||
errors.LogWarning(ctx, err.Error())
|
||||
}
|
||||
@@ -0,0 +1,147 @@
|
||||
package shadowsocks_2022
|
||||
|
||||
import (
|
||||
"context"
|
||||
"time"
|
||||
|
||||
shadowsocks "github.com/sagernet/sing-shadowsocks"
|
||||
"github.com/sagernet/sing-shadowsocks/shadowaead_2022"
|
||||
C "github.com/sagernet/sing/common"
|
||||
B "github.com/sagernet/sing/common/buf"
|
||||
"github.com/sagernet/sing/common/bufio"
|
||||
N "github.com/sagernet/sing/common/network"
|
||||
"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/session"
|
||||
"github.com/xtls/xray-core/common/signal"
|
||||
"github.com/xtls/xray-core/common/singbridge"
|
||||
"github.com/xtls/xray-core/transport"
|
||||
"github.com/xtls/xray-core/transport/internet"
|
||||
)
|
||||
|
||||
func init() {
|
||||
common.Must(common.RegisterConfig((*ClientConfig)(nil), func(ctx context.Context, config interface{}) (interface{}, error) {
|
||||
return NewClient(ctx, config.(*ClientConfig))
|
||||
}))
|
||||
}
|
||||
|
||||
type Outbound struct {
|
||||
ctx context.Context
|
||||
server net.Destination
|
||||
method shadowsocks.Method
|
||||
}
|
||||
|
||||
func NewClient(ctx context.Context, config *ClientConfig) (*Outbound, error) {
|
||||
o := &Outbound{
|
||||
ctx: ctx,
|
||||
server: net.Destination{
|
||||
Address: config.Address.AsAddress(),
|
||||
Port: net.Port(config.Port),
|
||||
Network: net.Network_TCP,
|
||||
},
|
||||
}
|
||||
if C.Contains(shadowaead_2022.List, config.Method) {
|
||||
if config.Key == "" {
|
||||
return nil, errors.New("missing psk")
|
||||
}
|
||||
method, err := shadowaead_2022.NewWithPassword(config.Method, config.Key, nil)
|
||||
if err != nil {
|
||||
return nil, errors.New("create method").Base(err)
|
||||
}
|
||||
o.method = method
|
||||
} else {
|
||||
return nil, errors.New("unknown method ", config.Method)
|
||||
}
|
||||
return o, nil
|
||||
}
|
||||
|
||||
func (o *Outbound) Process(ctx context.Context, link *transport.Link, dialer internet.Dialer) error {
|
||||
var inboundConn net.Conn
|
||||
inbound := session.InboundFromContext(ctx)
|
||||
if inbound != nil {
|
||||
inboundConn = inbound.Conn
|
||||
}
|
||||
|
||||
outbounds := session.OutboundsFromContext(ctx)
|
||||
ob := outbounds[len(outbounds)-1]
|
||||
if !ob.Target.IsValid() {
|
||||
return errors.New("target not specified")
|
||||
}
|
||||
ob.Name = "shadowsocks-2022"
|
||||
ob.CanSpliceCopy = 3
|
||||
destination := ob.Target
|
||||
network := destination.Network
|
||||
|
||||
errors.LogInfo(ctx, "tunneling request to ", destination, " via ", o.server.NetAddr())
|
||||
|
||||
serverDestination := o.server
|
||||
serverDestination.Network = network
|
||||
connection, err := dialer.Dial(ctx, serverDestination)
|
||||
if err != nil {
|
||||
return errors.New("failed to connect to server").Base(err)
|
||||
}
|
||||
defer connection.Close()
|
||||
|
||||
if session.TimeoutOnlyFromContext(ctx) {
|
||||
ctx, _ = context.WithCancel(context.Background())
|
||||
}
|
||||
|
||||
if network == net.Network_TCP {
|
||||
serverConn := o.method.DialEarlyConn(connection, singbridge.ToSocksaddr(destination))
|
||||
var handshake bool
|
||||
if timeoutReader, isTimeoutReader := link.Reader.(buf.TimeoutReader); isTimeoutReader {
|
||||
mb, err := timeoutReader.ReadMultiBufferTimeout(time.Millisecond * 100)
|
||||
if err != nil && err != buf.ErrNotTimeoutReader && err != buf.ErrReadTimeout {
|
||||
return errors.New("read payload").Base(err)
|
||||
}
|
||||
payload := B.New()
|
||||
for {
|
||||
payload.Reset()
|
||||
nb, n := buf.SplitBytes(mb, payload.FreeBytes())
|
||||
if n > 0 {
|
||||
payload.Truncate(n)
|
||||
_, err = serverConn.Write(payload.Bytes())
|
||||
if err != nil {
|
||||
payload.Release()
|
||||
return errors.New("write payload").Base(err)
|
||||
}
|
||||
handshake = true
|
||||
}
|
||||
if nb.IsEmpty() {
|
||||
break
|
||||
}
|
||||
mb = nb
|
||||
}
|
||||
payload.Release()
|
||||
}
|
||||
if !handshake {
|
||||
_, err = serverConn.Write(nil)
|
||||
if err != nil {
|
||||
return errors.New("client handshake").Base(err)
|
||||
}
|
||||
}
|
||||
return singbridge.CopyConn(ctx, inboundConn, link, serverConn)
|
||||
} else {
|
||||
var packetConn N.PacketConn
|
||||
if pc, isPacketConn := inboundConn.(N.PacketConn); isPacketConn {
|
||||
packetConn = pc
|
||||
} else if nc, isNetPacket := inboundConn.(net.PacketConn); isNetPacket {
|
||||
packetConn = bufio.NewPacketConn(nc)
|
||||
} else {
|
||||
packetConn = &singbridge.PacketConnWrapper{
|
||||
Reader: link.Reader,
|
||||
Writer: link.Writer,
|
||||
Conn: inboundConn,
|
||||
Dest: destination,
|
||||
T: signal.CancelAfterInactivity(ctx, func() {
|
||||
common.Interrupt(link.Reader)
|
||||
}, 300*time.Second),
|
||||
}
|
||||
}
|
||||
|
||||
serverConn := o.method.DialPacketConn(connection)
|
||||
return singbridge.ReturnError(bufio.CopyPacketConn(ctx, packetConn, serverConn))
|
||||
}
|
||||
}
|
||||
@@ -0,0 +1 @@
|
||||
package shadowsocks_2022
|
||||
+11
-2
@@ -32,6 +32,7 @@ type Config struct {
|
||||
AutoSystemRoutingTable []string `protobuf:"bytes,6,rep,name=auto_system_routing_table,json=autoSystemRoutingTable,proto3" json:"auto_system_routing_table,omitempty"`
|
||||
AutoOutboundsInterface string `protobuf:"bytes,7,opt,name=auto_outbounds_interface,json=autoOutboundsInterface,proto3" json:"auto_outbounds_interface,omitempty"`
|
||||
Desc string `protobuf:"bytes,8,opt,name=desc,proto3" json:"desc,omitempty"`
|
||||
Stack string `protobuf:"bytes,9,opt,name=stack,proto3" json:"stack,omitempty"`
|
||||
unknownFields protoimpl.UnknownFields
|
||||
sizeCache protoimpl.SizeCache
|
||||
}
|
||||
@@ -122,11 +123,18 @@ func (x *Config) GetDesc() string {
|
||||
return ""
|
||||
}
|
||||
|
||||
func (x *Config) GetStack() string {
|
||||
if x != nil {
|
||||
return x.Stack
|
||||
}
|
||||
return ""
|
||||
}
|
||||
|
||||
var File_proxy_tun_config_proto protoreflect.FileDescriptor
|
||||
|
||||
const file_proxy_tun_config_proto_rawDesc = "" +
|
||||
"\n" +
|
||||
"\x16proxy/tun/config.proto\x12\x0exray.proxy.tun\"\x82\x02\n" +
|
||||
"\x16proxy/tun/config.proto\x12\x0exray.proxy.tun\"\x98\x02\n" +
|
||||
"\x06Config\x12\x12\n" +
|
||||
"\x04name\x18\x01 \x01(\tR\x04name\x12\x10\n" +
|
||||
"\x03MTU\x18\x02 \x01(\rR\x03MTU\x12\x18\n" +
|
||||
@@ -136,7 +144,8 @@ const file_proxy_tun_config_proto_rawDesc = "" +
|
||||
"user_level\x18\x05 \x01(\rR\tuserLevel\x129\n" +
|
||||
"\x19auto_system_routing_table\x18\x06 \x03(\tR\x16autoSystemRoutingTable\x128\n" +
|
||||
"\x18auto_outbounds_interface\x18\a \x01(\tR\x16autoOutboundsInterface\x12\x12\n" +
|
||||
"\x04desc\x18\b \x01(\tR\x04descBL\n" +
|
||||
"\x04desc\x18\b \x01(\tR\x04desc\x12\x14\n" +
|
||||
"\x05stack\x18\t \x01(\tR\x05stackBL\n" +
|
||||
"\x12com.xray.proxy.tunP\x01Z#github.com/xtls/xray-core/proxy/tun\xaa\x02\x0eXray.Proxy.Tunb\x06proto3"
|
||||
|
||||
var (
|
||||
|
||||
@@ -15,4 +15,5 @@ message Config {
|
||||
repeated string auto_system_routing_table = 6;
|
||||
string auto_outbounds_interface = 7;
|
||||
string desc = 8;
|
||||
string stack = 9;
|
||||
}
|
||||
|
||||
+35
-4
@@ -37,6 +37,25 @@ type Handler struct {
|
||||
downlinkCounter stats.Counter
|
||||
}
|
||||
|
||||
type tunUDPStatsWriter struct {
|
||||
writer buf.Writer
|
||||
counter stats.Counter
|
||||
}
|
||||
|
||||
func (w *tunUDPStatsWriter) WriteMultiBuffer(mb buf.MultiBuffer) error {
|
||||
for len(mb) > 0 {
|
||||
remaining, packet := buf.SplitFirst(mb)
|
||||
packetSize := packet.Len()
|
||||
if err := w.writer.WriteMultiBuffer(buf.MultiBuffer{packet}); err != nil {
|
||||
buf.ReleaseMulti(remaining)
|
||||
return err
|
||||
}
|
||||
w.counter.Add(int64(packetSize))
|
||||
mb = remaining
|
||||
}
|
||||
return nil
|
||||
}
|
||||
|
||||
// ConnectionHandler interface with the only method that stack is going to push new connections to
|
||||
type ConnectionHandler interface {
|
||||
HandleConnection(conn net.Conn, destination net.Destination)
|
||||
@@ -104,7 +123,7 @@ func (t *Handler) Start() error {
|
||||
iface := updater.Get()
|
||||
if iface == nil {
|
||||
errors.LogInfo(context.Background(), "[tun] falied to set interface > iface == nil")
|
||||
return nil
|
||||
return errors.New("iface not found")
|
||||
}
|
||||
return c.Control(func(fd uintptr) {
|
||||
addrPort, _ := netip.ParseAddrPort(address)
|
||||
@@ -124,7 +143,9 @@ func (t *Handler) Start() error {
|
||||
|
||||
tunStackOptions := StackOptions{
|
||||
Tun: tunInterface,
|
||||
MTU: t.config.MTU,
|
||||
IdleTimeout: t.policyManager.ForLevel(t.config.UserLevel).Timeouts.ConnectionIdle,
|
||||
Backend: t.config.Stack,
|
||||
}
|
||||
tunStack, err := NewStack(t.ctx, tunStackOptions, t)
|
||||
if err != nil {
|
||||
@@ -171,7 +192,8 @@ func (t *Handler) HandleConnection(conn net.Conn, destination net.Destination) {
|
||||
return
|
||||
}
|
||||
source := net.DestinationFromAddr(remote)
|
||||
if t.uplinkCounter != nil || t.downlinkCounter != nil {
|
||||
isUDP := destination.Network == net.Network_UDP
|
||||
if !isUDP && (t.uplinkCounter != nil || t.downlinkCounter != nil) {
|
||||
conn = &stat.CounterConnection{
|
||||
Connection: conn,
|
||||
ReadCounter: t.uplinkCounter,
|
||||
@@ -203,9 +225,18 @@ func (t *Handler) HandleConnection(conn net.Conn, destination net.Destination) {
|
||||
})
|
||||
errors.LogInfo(ctx, "processing from ", source, " to ", destination)
|
||||
|
||||
reader := &buf.TimeoutWrapperReader{Reader: buf.NewReader(conn)}
|
||||
writer := buf.NewWriter(conn)
|
||||
if isUDP {
|
||||
reader.Counter = t.uplinkCounter
|
||||
if t.downlinkCounter != nil {
|
||||
writer = &tunUDPStatsWriter{writer: writer, counter: t.downlinkCounter}
|
||||
}
|
||||
}
|
||||
|
||||
link := &transport.Link{
|
||||
Reader: &buf.TimeoutWrapperReader{Reader: buf.NewReader(conn)},
|
||||
Writer: buf.NewWriter(conn),
|
||||
Reader: reader,
|
||||
Writer: writer,
|
||||
}
|
||||
if err := t.dispatcher.DispatchLink(ctx, destination, link); err != nil {
|
||||
errors.LogError(ctx, errors.New("connection closed").Base(err))
|
||||
|
||||
@@ -1,7 +1,10 @@
|
||||
package tun
|
||||
|
||||
import (
|
||||
"context"
|
||||
"time"
|
||||
|
||||
"github.com/xtls/xray-core/common/errors"
|
||||
)
|
||||
|
||||
// Stack interface implement ip protocol stack, bridging raw network packets and data streams
|
||||
@@ -13,5 +16,39 @@ type Stack interface {
|
||||
// StackOptions for the stack implementation
|
||||
type StackOptions struct {
|
||||
Tun Tun
|
||||
MTU uint32
|
||||
IdleTimeout time.Duration
|
||||
// Backend selects the concrete Stack implementation, see NewStack.
|
||||
Backend string
|
||||
}
|
||||
|
||||
const (
|
||||
// StackGVisor selects the full-featured gVisor based stack (default).
|
||||
StackGVisor = "gvisor"
|
||||
// StackSystem selects the lightweight, Xray-native stack, see newSystemStack.
|
||||
StackSystem = "system"
|
||||
)
|
||||
|
||||
// NewStack builds the ip stack selected by options.Backend.
|
||||
//
|
||||
// gVisor (the default/"gvisor" backend) is a general purpose stack, built
|
||||
// with the semantics needed for a real, lossy public network in mind:
|
||||
// congestion control, SACK/RACK loss recovery, retransmission timers, etc.
|
||||
// TUN traffic instead travels over a local, kernel-to-userspace channel that
|
||||
// neither reorders nor drops packets in normal operation, so none of that
|
||||
// complexity is actually required to shuffle bytes between it and the
|
||||
// dispatcher. The "system" backend trades gVisor's generality for a much
|
||||
// smaller, more direct code path tailored to that trusted, in-order channel:
|
||||
// no congestion control, no SACK/RACK, minimal buffering, and a plain RTO
|
||||
// timer as a safety net for the rare real loss, rather than a full
|
||||
// re-implementation of one. See stack_system.go for details.
|
||||
func NewStack(ctx context.Context, options StackOptions, handler *Handler) (Stack, error) {
|
||||
switch options.Backend {
|
||||
case "", StackGVisor:
|
||||
return newGVisorStack(ctx, options, handler)
|
||||
case StackSystem:
|
||||
return newSystemStack(ctx, options, handler)
|
||||
default:
|
||||
return nil, errors.New("unknown tun stack: ", options.Backend)
|
||||
}
|
||||
}
|
||||
|
||||
@@ -42,8 +42,8 @@ type stackGVisor struct {
|
||||
endpoint stack.LinkEndpoint
|
||||
}
|
||||
|
||||
// NewStack builds new ip stack (using gVisor)
|
||||
func NewStack(ctx context.Context, options StackOptions, handler *Handler) (Stack, error) {
|
||||
// newGVisorStack builds new ip stack (using gVisor)
|
||||
func newGVisorStack(ctx context.Context, options StackOptions, handler *Handler) (Stack, error) {
|
||||
gStack := &stackGVisor{
|
||||
ctx: ctx,
|
||||
tun: options.Tun,
|
||||
|
||||
@@ -0,0 +1,370 @@
|
||||
package tun
|
||||
|
||||
import (
|
||||
"context"
|
||||
"crypto/rand"
|
||||
"encoding/binary"
|
||||
"errors"
|
||||
"sync"
|
||||
"time"
|
||||
|
||||
xerrors "github.com/xtls/xray-core/common/errors"
|
||||
"github.com/xtls/xray-core/common/net"
|
||||
tunicmp "github.com/xtls/xray-core/proxy/tun/icmp"
|
||||
"gvisor.dev/gvisor/pkg/buffer"
|
||||
"gvisor.dev/gvisor/pkg/tcpip"
|
||||
"gvisor.dev/gvisor/pkg/tcpip/checksum"
|
||||
"gvisor.dev/gvisor/pkg/tcpip/header"
|
||||
"gvisor.dev/gvisor/pkg/tcpip/seqnum"
|
||||
"gvisor.dev/gvisor/pkg/tcpip/stack"
|
||||
)
|
||||
|
||||
// stackSystem is the lightweight, Xray-native ip stack, see NewStack.
|
||||
//
|
||||
// It reads and parses IPv4/IPv6 packets directly off the tun device (through
|
||||
// the GVisorDevice interface, already implemented for every supported
|
||||
// platform), without involving gVisor's stack.Stack, NIC or routing
|
||||
// machinery. UDP and ICMP echo reuse the exact same handlers as the gVisor
|
||||
// backend (udpConnectionHandler, tun/icmp) since those were already
|
||||
// implemented in terms of raw bytes. TCP is handled by a small dedicated
|
||||
// state machine, see stack_system_tcp.go.
|
||||
type stackSystem struct {
|
||||
ctx context.Context
|
||||
device GVisorDevice
|
||||
mtu uint32
|
||||
idleTimeout time.Duration
|
||||
// handler is stored as the narrower ConnectionHandler interface (which
|
||||
// *Handler satisfies) rather than *Handler itself, so the stack can be
|
||||
// exercised in tests with a lightweight fake, the same way stack_system_test.go does.
|
||||
handler ConnectionHandler
|
||||
|
||||
udp *udpConnectionHandler
|
||||
|
||||
tcpMu sync.Mutex
|
||||
tcp map[tcpKey]*tcpConn
|
||||
|
||||
cancel context.CancelFunc
|
||||
}
|
||||
|
||||
const systemStackDefaultMTU = 1500
|
||||
|
||||
// newSystemStack builds the lightweight "system" ip stack, see NewStack.
|
||||
func newSystemStack(ctx context.Context, options StackOptions, handler *Handler) (Stack, error) {
|
||||
device, ok := options.Tun.(GVisorDevice)
|
||||
if !ok {
|
||||
return nil, xerrors.New("tun stack \"system\" is not supported by this tun device")
|
||||
}
|
||||
mtu := options.MTU
|
||||
if mtu == 0 {
|
||||
mtu = systemStackDefaultMTU
|
||||
}
|
||||
return &stackSystem{
|
||||
ctx: ctx,
|
||||
device: device,
|
||||
mtu: mtu,
|
||||
idleTimeout: options.IdleTimeout,
|
||||
handler: handler,
|
||||
tcp: make(map[tcpKey]*tcpConn),
|
||||
}, nil
|
||||
}
|
||||
|
||||
// Start is called by Handler to bring the stack to life
|
||||
func (s *stackSystem) Start() error {
|
||||
ctx, cancel := context.WithCancel(s.ctx)
|
||||
s.cancel = cancel
|
||||
s.udp = newUdpConnectionHandler(s.handler.HandleConnection, s.writeRawUDPPacket)
|
||||
|
||||
go s.dispatchLoop(ctx)
|
||||
go s.idleReapLoop(ctx)
|
||||
return nil
|
||||
}
|
||||
|
||||
// Close is called by Handler to shut down the stack
|
||||
func (s *stackSystem) Close() error {
|
||||
if s.cancel != nil {
|
||||
s.cancel()
|
||||
}
|
||||
|
||||
s.tcpMu.Lock()
|
||||
conns := make([]*tcpConn, 0, len(s.tcp))
|
||||
for _, c := range s.tcp {
|
||||
conns = append(conns, c)
|
||||
}
|
||||
s.tcp = make(map[tcpKey]*tcpConn)
|
||||
s.tcpMu.Unlock()
|
||||
|
||||
for _, c := range conns {
|
||||
c.abort(errStackClosed)
|
||||
}
|
||||
|
||||
return nil
|
||||
}
|
||||
|
||||
// dispatchLoop reads and demultiplexes packets off the tun device, until ctx
|
||||
// is cancelled or the device fails permanently. It mirrors LinkEndpoint's own
|
||||
// dispatchLoop (stack_gvisor_endpoint.go), reusing the exact same GVisorDevice
|
||||
// contract, but hands packets to this file's own IPv4/IPv6 parsing instead of
|
||||
// gVisor's NIC/stack.Stack.
|
||||
func (s *stackSystem) dispatchLoop(ctx context.Context) {
|
||||
for {
|
||||
select {
|
||||
case <-ctx.Done():
|
||||
return
|
||||
default:
|
||||
}
|
||||
|
||||
version, packet, err := s.device.ReadPacket()
|
||||
if err != nil {
|
||||
if errors.Is(err, ErrQueueEmpty) {
|
||||
s.device.Wait()
|
||||
continue
|
||||
}
|
||||
return
|
||||
}
|
||||
|
||||
s.handlePacket(version, packet)
|
||||
packet.DecRef()
|
||||
}
|
||||
}
|
||||
|
||||
func (s *stackSystem) handlePacket(version byte, packet *stack.PacketBuffer) {
|
||||
data := concatSlices(packet.AsSlices())
|
||||
if len(data) == 0 {
|
||||
return
|
||||
}
|
||||
|
||||
switch version {
|
||||
case 4:
|
||||
s.handleIPv4(data)
|
||||
case 6:
|
||||
s.handleIPv6(data)
|
||||
}
|
||||
}
|
||||
|
||||
func concatSlices(slices [][]byte) []byte {
|
||||
if len(slices) == 1 {
|
||||
return slices[0]
|
||||
}
|
||||
total := 0
|
||||
for _, sl := range slices {
|
||||
total += len(sl)
|
||||
}
|
||||
if total == 0 {
|
||||
return nil
|
||||
}
|
||||
data := make([]byte, 0, total)
|
||||
for _, sl := range slices {
|
||||
data = append(data, sl...)
|
||||
}
|
||||
return data
|
||||
}
|
||||
|
||||
func (s *stackSystem) handleIPv4(data []byte) {
|
||||
hdr := header.IPv4(data)
|
||||
if !hdr.IsValid(len(data)) {
|
||||
return
|
||||
}
|
||||
// fragmentation is not supported: the tun MTU is expected to keep locally
|
||||
// generated packets from ever needing it, same as the gVisor backend's
|
||||
// default configuration
|
||||
if hdr.More() || hdr.FragmentOffset() != 0 {
|
||||
return
|
||||
}
|
||||
|
||||
s.handleTransport(header.IPv4ProtocolNumber, hdr.TransportProtocol(), hdr.SourceAddress(), hdr.DestinationAddress(), hdr.Payload())
|
||||
}
|
||||
|
||||
func (s *stackSystem) handleIPv6(data []byte) {
|
||||
hdr := header.IPv6(data)
|
||||
if !hdr.IsValid(len(data)) {
|
||||
return
|
||||
}
|
||||
|
||||
// only directly-encapsulated transport headers are handled, IPv6
|
||||
// extension headers (rare for ordinary locally generated traffic) are not
|
||||
// walked, same limitation as the fragmentation one above
|
||||
s.handleTransport(header.IPv6ProtocolNumber, hdr.TransportProtocol(), hdr.SourceAddress(), hdr.DestinationAddress(), hdr.Payload())
|
||||
}
|
||||
|
||||
func (s *stackSystem) handleTransport(netProto tcpip.NetworkProtocolNumber, transProto tcpip.TransportProtocolNumber, srcIP, dstIP tcpip.Address, payload []byte) {
|
||||
switch transProto {
|
||||
case header.TCPProtocolNumber:
|
||||
s.handleTCP(netProto, srcIP, dstIP, payload)
|
||||
case header.UDPProtocolNumber:
|
||||
s.handleUDP(netProto, srcIP, dstIP, payload)
|
||||
case header.ICMPv4ProtocolNumber:
|
||||
if netProto == header.IPv4ProtocolNumber {
|
||||
s.handleICMP(netProto, srcIP, dstIP, payload)
|
||||
}
|
||||
case header.ICMPv6ProtocolNumber:
|
||||
if netProto == header.IPv6ProtocolNumber {
|
||||
s.handleICMP(netProto, srcIP, dstIP, payload)
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
func (s *stackSystem) handleUDP(netProto tcpip.NetworkProtocolNumber, srcIP, dstIP tcpip.Address, payload []byte) {
|
||||
if len(payload) < header.UDPMinimumSize {
|
||||
return
|
||||
}
|
||||
udpHdr := header.UDP(payload)
|
||||
length := udpHdr.Length()
|
||||
if int(length) < header.UDPMinimumSize || int(length) > len(payload) {
|
||||
return
|
||||
}
|
||||
|
||||
// source/destination of the packet we process as incoming are, in other terms,
|
||||
// src is the side behind tun, dst is the side behind the dispatcher
|
||||
src := net.UDPDestination(net.IPAddress(srcIP.AsSlice()), net.Port(udpHdr.SourcePort()))
|
||||
dst := net.UDPDestination(net.IPAddress(dstIP.AsSlice()), net.Port(udpHdr.DestinationPort()))
|
||||
s.udp.HandlePacket(src, dst, payload[header.UDPMinimumSize:length])
|
||||
}
|
||||
|
||||
func (s *stackSystem) handleICMP(netProto tcpip.NetworkProtocolNumber, srcIP, dstIP tcpip.Address, message []byte) {
|
||||
ident, sequence, ok := tunicmp.ParseEchoRequest(netProto, message)
|
||||
if !ok {
|
||||
return
|
||||
}
|
||||
|
||||
reply, err := tunicmp.BuildLocalEchoReply(netProto, message, dstIP, srcIP)
|
||||
if err != nil {
|
||||
xerrors.LogInfoInner(s.ctx, err, "[tun] failed to build local icmp echo reply")
|
||||
return
|
||||
}
|
||||
|
||||
xerrors.LogDebug(s.ctx, "[tun][icmp] ", tunicmp.ProtocolLabel(netProto), " local echo reply ", dstIP, " -> ", srcIP, " id=", ident, " seq=", sequence)
|
||||
|
||||
transProto := header.ICMPv4ProtocolNumber
|
||||
if netProto == header.IPv6ProtocolNumber {
|
||||
transProto = header.ICMPv6ProtocolNumber
|
||||
}
|
||||
if err := s.writeTransportSegment(netProto, tcpip.TransportProtocolNumber(transProto), dstIP, srcIP, reply); err != nil {
|
||||
xerrors.LogInfoInner(s.ctx, err, "[tun] failed to write local icmp echo reply")
|
||||
}
|
||||
}
|
||||
|
||||
func (s *stackSystem) writeRawUDPPacket(payload []byte, src net.Destination, dst net.Destination) error {
|
||||
udpLen := header.UDPMinimumSize + len(payload)
|
||||
srcIP := tcpip.AddrFromSlice(src.Address.IP())
|
||||
dstIP := tcpip.AddrFromSlice(dst.Address.IP())
|
||||
|
||||
netProto := header.IPv4ProtocolNumber
|
||||
if !dst.Address.Family().IsIPv4() {
|
||||
netProto = header.IPv6ProtocolNumber
|
||||
}
|
||||
|
||||
segment := make([]byte, udpLen)
|
||||
udpHdr := header.UDP(segment)
|
||||
udpHdr.Encode(&header.UDPFields{
|
||||
SrcPort: uint16(src.Port),
|
||||
DstPort: uint16(dst.Port),
|
||||
Length: uint16(udpLen),
|
||||
})
|
||||
copy(segment[header.UDPMinimumSize:], payload)
|
||||
|
||||
xsum := header.PseudoHeaderChecksum(header.UDPProtocolNumber, srcIP, dstIP, uint16(udpLen))
|
||||
udpHdr.SetChecksum(^udpHdr.CalculateChecksum(checksum.Checksum(payload, xsum)))
|
||||
|
||||
return s.writeTransportSegment(netProto, header.UDPProtocolNumber, srcIP, dstIP, segment)
|
||||
}
|
||||
|
||||
// writeTransportSegment wraps a fully built, already checksummed transport
|
||||
// layer segment (UDP, ICMP or TCP) with an IP header and writes it to the tun
|
||||
// device.
|
||||
func (s *stackSystem) writeTransportSegment(netProto tcpip.NetworkProtocolNumber, transProto tcpip.TransportProtocolNumber, srcIP, dstIP tcpip.Address, segment []byte) error {
|
||||
ipHdrSize := header.IPv4MinimumSize
|
||||
if netProto == header.IPv6ProtocolNumber {
|
||||
ipHdrSize = header.IPv6MinimumSize
|
||||
}
|
||||
|
||||
pkt := stack.NewPacketBuffer(stack.PacketBufferOptions{
|
||||
ReserveHeaderBytes: ipHdrSize,
|
||||
Payload: buffer.MakeWithData(segment),
|
||||
})
|
||||
defer pkt.DecRef()
|
||||
|
||||
if netProto == header.IPv4ProtocolNumber {
|
||||
ipHdr := header.IPv4(pkt.NetworkHeader().Push(header.IPv4MinimumSize))
|
||||
ipHdr.Encode(&header.IPv4Fields{
|
||||
TotalLength: uint16(header.IPv4MinimumSize + len(segment)),
|
||||
TTL: 64,
|
||||
Protocol: uint8(transProto),
|
||||
SrcAddr: srcIP,
|
||||
DstAddr: dstIP,
|
||||
})
|
||||
ipHdr.SetChecksum(^ipHdr.CalculateChecksum())
|
||||
} else {
|
||||
ipHdr := header.IPv6(pkt.NetworkHeader().Push(header.IPv6MinimumSize))
|
||||
ipHdr.Encode(&header.IPv6Fields{
|
||||
PayloadLength: uint16(len(segment)),
|
||||
TransportProtocol: transProto,
|
||||
HopLimit: 64,
|
||||
SrcAddr: srcIP,
|
||||
DstAddr: dstIP,
|
||||
})
|
||||
}
|
||||
|
||||
if err := s.device.WritePacket(pkt); err != nil {
|
||||
return xerrors.New("failed to write raw packet: ", err.String())
|
||||
}
|
||||
return nil
|
||||
}
|
||||
|
||||
// idleReapLoop periodically aborts tcp connections that have seen no traffic
|
||||
// for longer than idleTimeout, finally putting that option to use (it was
|
||||
// tracked but never read anywhere before the "system" backend existed).
|
||||
func (s *stackSystem) idleReapLoop(ctx context.Context) {
|
||||
if s.idleTimeout <= 0 {
|
||||
return
|
||||
}
|
||||
interval := s.idleTimeout / 4
|
||||
if interval < time.Second {
|
||||
interval = time.Second
|
||||
}
|
||||
ticker := time.NewTicker(interval)
|
||||
defer ticker.Stop()
|
||||
|
||||
for {
|
||||
select {
|
||||
case <-ctx.Done():
|
||||
return
|
||||
case <-ticker.C:
|
||||
s.reapIdleConnections()
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
func (s *stackSystem) reapIdleConnections() {
|
||||
deadline := time.Now().Add(-s.idleTimeout)
|
||||
|
||||
s.tcpMu.Lock()
|
||||
var idle []*tcpConn
|
||||
for _, c := range s.tcp {
|
||||
if c.lastActiveTime().Before(deadline) {
|
||||
idle = append(idle, c)
|
||||
}
|
||||
}
|
||||
s.tcpMu.Unlock()
|
||||
|
||||
for _, c := range idle {
|
||||
c.abort(errConnIdleTimeout)
|
||||
}
|
||||
}
|
||||
|
||||
func (s *stackSystem) removeTCPConn(key tcpKey, c *tcpConn) {
|
||||
s.tcpMu.Lock()
|
||||
if existing, ok := s.tcp[key]; ok && existing == c {
|
||||
delete(s.tcp, key)
|
||||
}
|
||||
s.tcpMu.Unlock()
|
||||
}
|
||||
|
||||
// randomSequenceNumber returns a random initial sequence number for a new
|
||||
// connection. It doesn't need to be cryptographically unpredictable (the tun
|
||||
// channel is local and trusted), just varied enough to avoid confusion with
|
||||
// prior incarnations of the same 4-tuple.
|
||||
func randomSequenceNumber() seqnum.Value {
|
||||
var b [4]byte
|
||||
_, _ = rand.Read(b[:])
|
||||
return seqnum.Value(binary.BigEndian.Uint32(b[:]))
|
||||
}
|
||||
@@ -0,0 +1,725 @@
|
||||
package tun
|
||||
|
||||
import (
|
||||
"io"
|
||||
"sync"
|
||||
"time"
|
||||
|
||||
xerrors "github.com/xtls/xray-core/common/errors"
|
||||
"github.com/xtls/xray-core/common/net"
|
||||
"gvisor.dev/gvisor/pkg/tcpip"
|
||||
"gvisor.dev/gvisor/pkg/tcpip/checksum"
|
||||
"gvisor.dev/gvisor/pkg/tcpip/header"
|
||||
"gvisor.dev/gvisor/pkg/tcpip/seqnum"
|
||||
)
|
||||
|
||||
// This file implements a small, dedicated TCP state machine for the "system"
|
||||
// tun stack, see stack_system.go. It intentionally does not implement window
|
||||
// scaling, SACK, timestamps, congestion control, fast retransmit or
|
||||
// out-of-order reassembly: the tun channel only ever carries packets produced
|
||||
// by the local OS network stack and handed to us directly, so it neither
|
||||
// reorders nor drops them the way the public internet does; a single RTO
|
||||
// timer (also used for zero-window probing) is enough to make the connection
|
||||
// robust against the rare occasions a segment does not make it through.
|
||||
const (
|
||||
minRTO = 300 * time.Millisecond
|
||||
maxRTO = 30 * time.Second
|
||||
maxRTORetries = 12
|
||||
lingerDuration = 5 * time.Second
|
||||
|
||||
// maxSendBuffer/maxRecvBuffer match the gVisor backend's own default
|
||||
// buffer sizes (tcp.DefaultSendBufferSize/DefaultReceiveBufferSize), so
|
||||
// switching between backends does not change buffering expectations.
|
||||
maxSendBuffer = 1 << 20
|
||||
maxRecvBuffer = 1 << 20
|
||||
)
|
||||
|
||||
var (
|
||||
errStackClosed = xerrors.New("tun stack closed")
|
||||
errConnReset = xerrors.New("connection reset by peer")
|
||||
errConnClosed = xerrors.New("use of closed network connection")
|
||||
errConnIdleTimeout = xerrors.New("connection idle timeout")
|
||||
errConnTimedOut = xerrors.New("connection timed out")
|
||||
)
|
||||
|
||||
type tcpState uint8
|
||||
|
||||
const (
|
||||
stateSynRcvd tcpState = iota
|
||||
stateEstablished
|
||||
stateCloseWait // peer's FIN was received; we may still send until we close too
|
||||
stateClosing // our FIN was sent (from Established or CloseWait)
|
||||
stateTimeWait // both FINs exchanged and acked; short linger before removal
|
||||
stateClosed // terminal, removed from the connection table
|
||||
)
|
||||
|
||||
// tcpKey identifies a tcp connection the same way it appears on the wire
|
||||
// flowing from the app behind the tun device towards its destination.
|
||||
type tcpKey struct {
|
||||
netProto tcpip.NetworkProtocolNumber
|
||||
srcAddr tcpip.Address
|
||||
srcPort uint16
|
||||
dstAddr tcpip.Address
|
||||
dstPort uint16
|
||||
}
|
||||
|
||||
// tcpConn is a minimal TCP endpoint implementing net.Conn. It deliberately
|
||||
// exposes only plain Read/Write (never ReadMultiBuffer/WriteMultiBuffer) so
|
||||
// that stat.CounterConnection in handler.go keeps accounting traffic
|
||||
// correctly, matching the udpConn precedent in udp_fullcone.go.
|
||||
type tcpConn struct {
|
||||
stack *stackSystem
|
||||
key tcpKey
|
||||
src net.Destination
|
||||
dst net.Destination
|
||||
|
||||
ourMSS int
|
||||
|
||||
mu sync.Mutex
|
||||
cond *sync.Cond
|
||||
|
||||
state tcpState
|
||||
|
||||
// send side. sendQueue[0] always holds the byte at sequence sndUna: acked
|
||||
// bytes are trimmed off the front, so no separate "acked" bookkeeping is
|
||||
// needed. sendQueue[:unsentOffset] has been transmitted at least once;
|
||||
// sendQueue[unsentOffset:] never has.
|
||||
iss seqnum.Value
|
||||
sndUna seqnum.Value
|
||||
sndNxt seqnum.Value
|
||||
sndMSS int
|
||||
peerWindow uint32
|
||||
sendQueue []byte
|
||||
unsentOffset int
|
||||
closeCalled bool
|
||||
finSent bool
|
||||
finAcked bool
|
||||
finSeq seqnum.Value
|
||||
|
||||
// receive side.
|
||||
irs seqnum.Value
|
||||
rcvNxt seqnum.Value
|
||||
recvQueue [][]byte
|
||||
recvOffset int
|
||||
recvBuffered int
|
||||
recvClosed bool
|
||||
|
||||
err error
|
||||
|
||||
lastActive time.Time
|
||||
|
||||
rtoTimer *time.Timer
|
||||
rtoBackoff int
|
||||
lingerTimer *time.Timer
|
||||
}
|
||||
|
||||
var _ net.Conn = (*tcpConn)(nil)
|
||||
|
||||
// outgoingMSS returns the MSS we can use without ever needing IP
|
||||
// fragmentation (unsupported), given the tun device's MTU.
|
||||
func outgoingMSS(mtu uint32, netProto tcpip.NetworkProtocolNumber) int {
|
||||
ipHdrSize := header.IPv4MinimumSize
|
||||
if netProto == header.IPv6ProtocolNumber {
|
||||
ipHdrSize = header.IPv6MinimumSize
|
||||
}
|
||||
mss := int(mtu) - ipHdrSize - header.TCPMinimumSize
|
||||
const minMSS = 88
|
||||
if mss < minMSS {
|
||||
mss = minMSS
|
||||
}
|
||||
return mss
|
||||
}
|
||||
|
||||
// handleTCP is the tcp entry point from stackSystem.handleTransport.
|
||||
func (s *stackSystem) handleTCP(netProto tcpip.NetworkProtocolNumber, srcIP, dstIP tcpip.Address, payload []byte) {
|
||||
if len(payload) < header.TCPMinimumSize {
|
||||
return
|
||||
}
|
||||
tcpHdr := header.TCP(payload)
|
||||
if _, _, ok := header.TCPValid(tcpHdr, nil, 0, tcpip.Address{}, tcpip.Address{}, true); !ok {
|
||||
return
|
||||
}
|
||||
|
||||
key := tcpKey{
|
||||
netProto: netProto,
|
||||
srcAddr: srcIP,
|
||||
srcPort: tcpHdr.SourcePort(),
|
||||
dstAddr: dstIP,
|
||||
dstPort: tcpHdr.DestinationPort(),
|
||||
}
|
||||
|
||||
s.tcpMu.Lock()
|
||||
conn, ok := s.tcp[key]
|
||||
s.tcpMu.Unlock()
|
||||
|
||||
if ok {
|
||||
conn.handleSegment(tcpHdr)
|
||||
return
|
||||
}
|
||||
|
||||
flags := tcpHdr.Flags()
|
||||
if flags&header.TCPFlagRst != 0 {
|
||||
return // never generate a reset in response to a reset
|
||||
}
|
||||
if flags&header.TCPFlagSyn != 0 && flags&header.TCPFlagAck == 0 {
|
||||
s.newTCPConn(key, tcpHdr)
|
||||
return
|
||||
}
|
||||
|
||||
// any other segment referencing an unknown connection: let the peer know
|
||||
// promptly it no longer/never existed, same as a real kernel would
|
||||
s.sendRawTCPReset(key, tcpHdr)
|
||||
}
|
||||
|
||||
func (s *stackSystem) newTCPConn(key tcpKey, tcpHdr header.TCP) {
|
||||
synOpts := header.ParseSynOptions(tcpHdr.Options(), false)
|
||||
|
||||
c := &tcpConn{
|
||||
stack: s,
|
||||
key: key,
|
||||
src: net.TCPDestination(net.IPAddress(key.srcAddr.AsSlice()), net.Port(key.srcPort)),
|
||||
dst: net.TCPDestination(net.IPAddress(key.dstAddr.AsSlice()), net.Port(key.dstPort)),
|
||||
state: stateSynRcvd,
|
||||
}
|
||||
c.cond = sync.NewCond(&c.mu)
|
||||
|
||||
c.iss = randomSequenceNumber()
|
||||
c.sndUna = c.iss
|
||||
c.sndNxt = c.iss.Add(1)
|
||||
|
||||
c.irs = seqnum.Value(tcpHdr.SequenceNumber())
|
||||
c.rcvNxt = c.irs.Add(1)
|
||||
|
||||
c.ourMSS = outgoingMSS(s.mtu, key.netProto)
|
||||
c.sndMSS = int(synOpts.MSS)
|
||||
if c.sndMSS <= 0 || c.sndMSS > c.ourMSS {
|
||||
c.sndMSS = c.ourMSS
|
||||
}
|
||||
c.lastActive = time.Now()
|
||||
|
||||
s.tcpMu.Lock()
|
||||
s.tcp[key] = c
|
||||
s.tcpMu.Unlock()
|
||||
|
||||
c.mu.Lock()
|
||||
c.sendSynAckLocked()
|
||||
c.mu.Unlock()
|
||||
}
|
||||
|
||||
// sendRawTCPReset replies to a segment that does not match any known
|
||||
// connection, following the rules of RFC 9293 §3.10.7.1.
|
||||
func (s *stackSystem) sendRawTCPReset(key tcpKey, tcpHdr header.TCP) {
|
||||
flags := tcpHdr.Flags()
|
||||
segLen := seqnum.Size(len(tcpHdr.Payload()))
|
||||
if flags&header.TCPFlagSyn != 0 {
|
||||
segLen++
|
||||
}
|
||||
if flags&header.TCPFlagFin != 0 {
|
||||
segLen++
|
||||
}
|
||||
|
||||
var seq, ack seqnum.Value
|
||||
var ackFlag header.TCPFlags
|
||||
if flags&header.TCPFlagAck != 0 {
|
||||
seq = seqnum.Value(tcpHdr.AckNumber())
|
||||
} else {
|
||||
ack = seqnum.Value(tcpHdr.SequenceNumber()).Add(segLen)
|
||||
ackFlag = header.TCPFlagAck
|
||||
}
|
||||
|
||||
segment := make([]byte, header.TCPMinimumSize)
|
||||
rst := header.TCP(segment)
|
||||
rst.Encode(&header.TCPFields{
|
||||
SrcPort: key.dstPort,
|
||||
DstPort: key.srcPort,
|
||||
SeqNum: uint32(seq),
|
||||
AckNum: uint32(ack),
|
||||
DataOffset: header.TCPMinimumSize,
|
||||
Flags: header.TCPFlagRst | ackFlag,
|
||||
WindowSize: 0,
|
||||
})
|
||||
xsum := header.PseudoHeaderChecksum(header.TCPProtocolNumber, key.dstAddr, key.srcAddr, uint16(len(segment)))
|
||||
rst.SetChecksum(^rst.CalculateChecksum(xsum))
|
||||
|
||||
if err := s.writeTransportSegment(key.netProto, header.TCPProtocolNumber, key.dstAddr, key.srcAddr, segment); err != nil {
|
||||
xerrors.LogInfoInner(s.ctx, err, "[tun] failed to write tcp reset")
|
||||
}
|
||||
}
|
||||
|
||||
func (c *tcpConn) lastActiveTime() time.Time {
|
||||
c.mu.Lock()
|
||||
defer c.mu.Unlock()
|
||||
return c.lastActive
|
||||
}
|
||||
|
||||
// abort is the externally callable (unlocked) equivalent of abortLocked,
|
||||
// used by the idle reaper and by Close's callers indirectly through it.
|
||||
func (c *tcpConn) abort(err error) {
|
||||
c.mu.Lock()
|
||||
defer c.mu.Unlock()
|
||||
c.abortLocked(err)
|
||||
}
|
||||
|
||||
func (c *tcpConn) abortLocked(err error) {
|
||||
if c.state == stateClosed {
|
||||
return
|
||||
}
|
||||
c.state = stateClosed
|
||||
c.stopRTOLocked()
|
||||
c.stopLingerLocked()
|
||||
c.err = err
|
||||
c.cond.Broadcast()
|
||||
// deliberately does not send an RST: if the peer sends anything else for
|
||||
// this connection later, it will miss the (now removed) table entry and
|
||||
// get a fresh, correctly-addressed reset from sendRawTCPReset above.
|
||||
c.stack.removeTCPConn(c.key, c)
|
||||
}
|
||||
|
||||
// handleSegment processes one already-demultiplexed incoming segment.
|
||||
func (c *tcpConn) handleSegment(tcpHdr header.TCP) {
|
||||
c.mu.Lock()
|
||||
defer c.mu.Unlock()
|
||||
|
||||
if c.state == stateClosed {
|
||||
return
|
||||
}
|
||||
c.lastActive = time.Now()
|
||||
|
||||
flags := tcpHdr.Flags()
|
||||
|
||||
if flags&header.TCPFlagRst != 0 {
|
||||
c.abortLocked(errConnReset)
|
||||
return
|
||||
}
|
||||
|
||||
if c.state == stateSynRcvd {
|
||||
c.handleSynRcvdSegmentLocked(tcpHdr)
|
||||
return
|
||||
}
|
||||
|
||||
if flags&header.TCPFlagSyn != 0 {
|
||||
// unexpected SYN on an already-established connection is not
|
||||
// modeled; treat it like the peer abandoned and reset it
|
||||
c.abortLocked(errConnReset)
|
||||
return
|
||||
}
|
||||
|
||||
if flags&header.TCPFlagAck != 0 {
|
||||
c.handleAckLocked(seqnum.Value(tcpHdr.AckNumber()), tcpHdr.WindowSize())
|
||||
}
|
||||
|
||||
c.acceptInOrderLocked(seqnum.Value(tcpHdr.SequenceNumber()), tcpHdr.Payload(), flags&header.TCPFlagFin != 0)
|
||||
}
|
||||
|
||||
func (c *tcpConn) handleSynRcvdSegmentLocked(tcpHdr header.TCP) {
|
||||
flags := tcpHdr.Flags()
|
||||
|
||||
if flags&header.TCPFlagSyn != 0 {
|
||||
// peer's retransmission of the original SYN, our SYN-ACK likely
|
||||
// hasn't reached them yet: resend it and rely entirely on their own
|
||||
// retransmission timer rather than running one on our side too
|
||||
c.sendSynAckLocked()
|
||||
return
|
||||
}
|
||||
if flags&header.TCPFlagAck == 0 {
|
||||
return
|
||||
}
|
||||
if seqnum.Value(tcpHdr.AckNumber()) != c.sndNxt {
|
||||
// does not acknowledge our SYN correctly; a well-behaved peer will
|
||||
// simply retry, so it is safe to just ignore this segment
|
||||
return
|
||||
}
|
||||
|
||||
c.state = stateEstablished
|
||||
go c.stack.handler.HandleConnection(c, c.dst)
|
||||
|
||||
c.handleAckLocked(seqnum.Value(tcpHdr.AckNumber()), tcpHdr.WindowSize())
|
||||
c.acceptInOrderLocked(seqnum.Value(tcpHdr.SequenceNumber()), tcpHdr.Payload(), flags&header.TCPFlagFin != 0)
|
||||
}
|
||||
|
||||
// acceptInOrderLocked handles the data/FIN portion of a segment once it is
|
||||
// known to be neither a SYN nor a RST. Only strictly in-order segments are
|
||||
// accepted; anything else is dropped (relying on the peer's retransmission)
|
||||
// since the tun channel is expected to already deliver packets in order.
|
||||
func (c *tcpConn) acceptInOrderLocked(seq seqnum.Value, payload []byte, fin bool) {
|
||||
if seq != c.rcvNxt {
|
||||
c.sendAckLocked()
|
||||
return
|
||||
}
|
||||
|
||||
accept := payload
|
||||
if room := c.recvWindowLocked(); uint32(len(accept)) > room {
|
||||
accept = accept[:room]
|
||||
}
|
||||
if len(accept) > 0 {
|
||||
c.enqueueRecvLocked(accept)
|
||||
c.rcvNxt = c.rcvNxt.Add(seqnum.Size(len(accept)))
|
||||
}
|
||||
|
||||
finAccepted := false
|
||||
if fin && len(accept) == len(payload) {
|
||||
c.onFinLocked()
|
||||
c.rcvNxt = c.rcvNxt.Add(1)
|
||||
finAccepted = true
|
||||
}
|
||||
|
||||
if len(accept) > 0 || finAccepted || len(accept) < len(payload) {
|
||||
c.sendAckLocked()
|
||||
}
|
||||
}
|
||||
|
||||
func (c *tcpConn) onFinLocked() {
|
||||
if c.recvClosed {
|
||||
return
|
||||
}
|
||||
c.recvClosed = true
|
||||
c.cond.Broadcast()
|
||||
if c.state == stateEstablished {
|
||||
c.state = stateCloseWait
|
||||
}
|
||||
c.maybeFinishCloseLocked()
|
||||
}
|
||||
|
||||
func (c *tcpConn) handleAckLocked(ackNum seqnum.Value, windowSize uint16) {
|
||||
if ackNum.LessThan(c.sndUna) {
|
||||
// old/duplicate ack: no fast-retransmit heuristics implemented
|
||||
c.peerWindow = uint32(windowSize)
|
||||
c.trySendLocked()
|
||||
return
|
||||
}
|
||||
if c.sndNxt.LessThan(ackNum) {
|
||||
// acknowledges more than we ever sent: lenient clamp instead of
|
||||
// rejecting the segment outright
|
||||
ackNum = c.sndNxt
|
||||
}
|
||||
|
||||
if advanced := c.sndUna.Size(ackNum); advanced > 0 {
|
||||
c.sndUna = ackNum
|
||||
n := int(advanced)
|
||||
if n > len(c.sendQueue) {
|
||||
n = len(c.sendQueue)
|
||||
}
|
||||
c.sendQueue = c.sendQueue[n:]
|
||||
c.unsentOffset -= n
|
||||
if c.unsentOffset < 0 {
|
||||
c.unsentOffset = 0
|
||||
}
|
||||
c.rtoBackoff = 0
|
||||
if c.finSent && c.sndUna == c.sndNxt {
|
||||
c.finAcked = true
|
||||
}
|
||||
c.cond.Broadcast()
|
||||
}
|
||||
|
||||
c.peerWindow = uint32(windowSize)
|
||||
c.trySendLocked()
|
||||
c.maybeFinishCloseLocked()
|
||||
}
|
||||
|
||||
func (c *tcpConn) maybeFinishCloseLocked() {
|
||||
if c.state == stateClosing && c.finAcked && c.recvClosed {
|
||||
c.state = stateTimeWait
|
||||
c.startLingerLocked()
|
||||
}
|
||||
}
|
||||
|
||||
func (c *tcpConn) recvWindowLocked() uint32 {
|
||||
room := maxRecvBuffer - c.recvBuffered
|
||||
if room < 0 {
|
||||
room = 0
|
||||
}
|
||||
if room > 0xffff {
|
||||
room = 0xffff
|
||||
}
|
||||
return uint32(room)
|
||||
}
|
||||
|
||||
func (c *tcpConn) enqueueRecvLocked(payload []byte) {
|
||||
data := make([]byte, len(payload))
|
||||
copy(data, payload)
|
||||
c.recvQueue = append(c.recvQueue, data)
|
||||
c.recvBuffered += len(data)
|
||||
c.cond.Broadcast()
|
||||
}
|
||||
|
||||
// sendOneChunkLocked transmits up to maxLen bytes of never-yet-sent data (if
|
||||
// any remains), advancing sndNxt/unsentOffset. It returns the number of
|
||||
// bytes sent, 0 if none remained.
|
||||
func (c *tcpConn) sendOneChunkLocked(maxLen int) int {
|
||||
remaining := len(c.sendQueue) - c.unsentOffset
|
||||
if remaining <= 0 {
|
||||
return 0
|
||||
}
|
||||
if maxLen > remaining {
|
||||
maxLen = remaining
|
||||
}
|
||||
if maxLen > c.sndMSS {
|
||||
maxLen = c.sndMSS
|
||||
}
|
||||
if maxLen <= 0 {
|
||||
return 0
|
||||
}
|
||||
data := c.sendQueue[c.unsentOffset : c.unsentOffset+maxLen]
|
||||
c.sendDataSegmentLocked(c.sndNxt, data, false)
|
||||
c.sndNxt = c.sndNxt.Add(seqnum.Size(maxLen))
|
||||
c.unsentOffset += maxLen
|
||||
return maxLen
|
||||
}
|
||||
|
||||
func (c *tcpConn) trySendLocked() {
|
||||
switch c.state {
|
||||
case stateSynRcvd, stateTimeWait, stateClosed:
|
||||
return
|
||||
}
|
||||
|
||||
for {
|
||||
inFlight := int(c.sndUna.Size(c.sndNxt))
|
||||
windowLeft := int(c.peerWindow) - inFlight
|
||||
if windowLeft <= 0 {
|
||||
break
|
||||
}
|
||||
if c.sendOneChunkLocked(windowLeft) == 0 {
|
||||
break
|
||||
}
|
||||
}
|
||||
|
||||
if c.closeCalled && !c.finSent && c.unsentOffset == len(c.sendQueue) {
|
||||
c.finSeq = c.sndNxt
|
||||
c.sendDataSegmentLocked(c.finSeq, nil, true)
|
||||
c.sndNxt = c.sndNxt.Add(1)
|
||||
c.finSent = true
|
||||
}
|
||||
|
||||
c.refreshRTOLocked()
|
||||
}
|
||||
|
||||
func (c *tcpConn) outstandingLocked() bool {
|
||||
if c.unsentOffset > 0 {
|
||||
return true // already-transmitted data pending ack
|
||||
}
|
||||
if len(c.sendQueue) > c.unsentOffset && c.peerWindow == 0 {
|
||||
return true // blocked purely by a zero window; need to probe
|
||||
}
|
||||
if c.finSent && !c.finAcked {
|
||||
return true // FIN transmitted but not yet acked
|
||||
}
|
||||
return false
|
||||
}
|
||||
|
||||
func (c *tcpConn) refreshRTOLocked() {
|
||||
if c.outstandingLocked() {
|
||||
c.scheduleRTOLocked()
|
||||
} else {
|
||||
c.stopRTOLocked()
|
||||
}
|
||||
}
|
||||
|
||||
func (c *tcpConn) rtoDurationLocked() time.Duration {
|
||||
d := minRTO * time.Duration(uint64(1)<<uint(c.rtoBackoff))
|
||||
if d > maxRTO || d <= 0 {
|
||||
d = maxRTO
|
||||
}
|
||||
return d
|
||||
}
|
||||
|
||||
func (c *tcpConn) scheduleRTOLocked() {
|
||||
d := c.rtoDurationLocked()
|
||||
if c.rtoTimer == nil {
|
||||
c.rtoTimer = time.AfterFunc(d, c.onRTOTimerFired)
|
||||
} else {
|
||||
c.rtoTimer.Reset(d)
|
||||
}
|
||||
}
|
||||
|
||||
func (c *tcpConn) stopRTOLocked() {
|
||||
if c.rtoTimer != nil {
|
||||
c.rtoTimer.Stop()
|
||||
}
|
||||
}
|
||||
|
||||
func (c *tcpConn) onRTOTimerFired() {
|
||||
c.mu.Lock()
|
||||
defer c.mu.Unlock()
|
||||
c.onRTOFireLocked()
|
||||
}
|
||||
|
||||
func (c *tcpConn) onRTOFireLocked() {
|
||||
if c.state == stateClosed || !c.outstandingLocked() {
|
||||
return
|
||||
}
|
||||
if c.rtoBackoff >= maxRTORetries {
|
||||
c.abortLocked(errConnTimedOut)
|
||||
return
|
||||
}
|
||||
c.rtoBackoff++
|
||||
|
||||
switch {
|
||||
case c.unsentOffset > 0:
|
||||
c.sendDataSegmentLocked(c.sndUna, c.sendQueue[:c.unsentOffset], false)
|
||||
case len(c.sendQueue) > c.unsentOffset:
|
||||
// nothing in flight, but blocked by a zero peer window: probe with
|
||||
// exactly one new byte, per RFC 9293 §3.8.6.1
|
||||
c.sendOneChunkLocked(1)
|
||||
case c.finSent && !c.finAcked:
|
||||
c.sendDataSegmentLocked(c.finSeq, nil, true)
|
||||
}
|
||||
|
||||
c.refreshRTOLocked()
|
||||
}
|
||||
|
||||
func (c *tcpConn) startLingerLocked() {
|
||||
c.stopRTOLocked()
|
||||
c.lingerTimer = time.AfterFunc(lingerDuration, func() {
|
||||
c.mu.Lock()
|
||||
defer c.mu.Unlock()
|
||||
c.abortLocked(errConnClosed)
|
||||
})
|
||||
}
|
||||
|
||||
func (c *tcpConn) stopLingerLocked() {
|
||||
if c.lingerTimer != nil {
|
||||
c.lingerTimer.Stop()
|
||||
}
|
||||
}
|
||||
|
||||
// transmitLocked builds, checksums and writes a single tcp segment.
|
||||
func (c *tcpConn) transmitLocked(seq, ack seqnum.Value, flags header.TCPFlags, payload []byte, options []byte) {
|
||||
headerLen := header.TCPMinimumSize + len(options)
|
||||
segment := make([]byte, headerLen+len(payload))
|
||||
tcpHdr := header.TCP(segment)
|
||||
tcpHdr.Encode(&header.TCPFields{
|
||||
SrcPort: c.key.dstPort,
|
||||
DstPort: c.key.srcPort,
|
||||
SeqNum: uint32(seq),
|
||||
AckNum: uint32(ack),
|
||||
DataOffset: uint8(headerLen),
|
||||
Flags: flags,
|
||||
WindowSize: uint16(c.recvWindowLocked()),
|
||||
})
|
||||
copy(tcpHdr.Options(), options)
|
||||
copy(segment[headerLen:], payload)
|
||||
|
||||
xsum := header.PseudoHeaderChecksum(header.TCPProtocolNumber, c.key.dstAddr, c.key.srcAddr, uint16(len(segment)))
|
||||
xsum = checksum.Checksum(payload, xsum)
|
||||
tcpHdr.SetChecksum(^tcpHdr.CalculateChecksum(xsum))
|
||||
|
||||
if err := c.stack.writeTransportSegment(c.key.netProto, header.TCPProtocolNumber, c.key.dstAddr, c.key.srcAddr, segment); err != nil {
|
||||
xerrors.LogInfoInner(c.stack.ctx, err, "[tun] failed to write tcp segment")
|
||||
}
|
||||
}
|
||||
|
||||
func (c *tcpConn) sendSynAckLocked() {
|
||||
var optBuf [header.TCPOptionMSSLength]byte
|
||||
n := header.EncodeMSSOption(uint32(c.ourMSS), optBuf[:])
|
||||
c.transmitLocked(c.iss, c.rcvNxt, header.TCPFlagSyn|header.TCPFlagAck, nil, optBuf[:n])
|
||||
}
|
||||
|
||||
func (c *tcpConn) sendAckLocked() {
|
||||
c.transmitLocked(c.sndNxt, c.rcvNxt, header.TCPFlagAck, nil, nil)
|
||||
}
|
||||
|
||||
func (c *tcpConn) sendDataSegmentLocked(seq seqnum.Value, payload []byte, fin bool) {
|
||||
flags := header.TCPFlagAck
|
||||
if fin {
|
||||
flags |= header.TCPFlagFin
|
||||
}
|
||||
c.transmitLocked(seq, c.rcvNxt, flags, payload, nil)
|
||||
}
|
||||
|
||||
// Read implements net.Conn.
|
||||
func (c *tcpConn) Read(p []byte) (int, error) {
|
||||
c.mu.Lock()
|
||||
defer c.mu.Unlock()
|
||||
|
||||
for len(c.recvQueue) == 0 && c.err == nil && !c.recvClosed {
|
||||
c.cond.Wait()
|
||||
}
|
||||
if c.err != nil {
|
||||
return 0, c.err
|
||||
}
|
||||
if len(c.recvQueue) == 0 {
|
||||
return 0, io.EOF
|
||||
}
|
||||
|
||||
before := c.recvWindowLocked()
|
||||
|
||||
chunk := c.recvQueue[0]
|
||||
n := copy(p, chunk[c.recvOffset:])
|
||||
c.recvOffset += n
|
||||
c.recvBuffered -= n
|
||||
if c.recvOffset == len(chunk) {
|
||||
c.recvQueue = c.recvQueue[1:]
|
||||
c.recvOffset = 0
|
||||
}
|
||||
|
||||
// let the peer know promptly if reading just freed up a previously
|
||||
// exhausted window, instead of waiting for it to probe us for an update
|
||||
if after := c.recvWindowLocked(); before == 0 && after > 0 {
|
||||
c.sendAckLocked()
|
||||
}
|
||||
|
||||
return n, nil
|
||||
}
|
||||
|
||||
// Write implements net.Conn.
|
||||
func (c *tcpConn) Write(p []byte) (int, error) {
|
||||
c.mu.Lock()
|
||||
defer c.mu.Unlock()
|
||||
|
||||
if c.closeCalled {
|
||||
return 0, errConnClosed
|
||||
}
|
||||
|
||||
total := 0
|
||||
for total < len(p) {
|
||||
if c.err != nil {
|
||||
return total, c.err
|
||||
}
|
||||
if c.closeCalled {
|
||||
return total, errConnClosed
|
||||
}
|
||||
room := maxSendBuffer - len(c.sendQueue)
|
||||
if room <= 0 {
|
||||
c.cond.Wait()
|
||||
continue
|
||||
}
|
||||
n := len(p) - total
|
||||
if n > room {
|
||||
n = room
|
||||
}
|
||||
c.sendQueue = append(c.sendQueue, p[total:total+n]...)
|
||||
total += n
|
||||
}
|
||||
|
||||
c.trySendLocked()
|
||||
return total, nil
|
||||
}
|
||||
|
||||
// Close implements net.Conn.
|
||||
func (c *tcpConn) Close() error {
|
||||
c.mu.Lock()
|
||||
defer c.mu.Unlock()
|
||||
|
||||
if c.closeCalled {
|
||||
return nil
|
||||
}
|
||||
c.closeCalled = true
|
||||
c.cond.Broadcast()
|
||||
|
||||
switch c.state {
|
||||
case stateEstablished, stateCloseWait:
|
||||
c.state = stateClosing
|
||||
default:
|
||||
return nil
|
||||
}
|
||||
|
||||
c.trySendLocked()
|
||||
return nil
|
||||
}
|
||||
|
||||
func (c *tcpConn) LocalAddr() net.Addr { return c.dst.RawNetAddr() }
|
||||
func (c *tcpConn) RemoteAddr() net.Addr { return c.src.RawNetAddr() }
|
||||
|
||||
func (c *tcpConn) SetDeadline(t time.Time) error { return nil }
|
||||
func (c *tcpConn) SetReadDeadline(t time.Time) error { return nil }
|
||||
func (c *tcpConn) SetWriteDeadline(t time.Time) error { return nil }
|
||||
@@ -0,0 +1,448 @@
|
||||
package tun
|
||||
|
||||
import (
|
||||
"context"
|
||||
"io"
|
||||
"sync"
|
||||
"testing"
|
||||
"time"
|
||||
|
||||
"github.com/xtls/xray-core/common/net"
|
||||
"gvisor.dev/gvisor/pkg/buffer"
|
||||
"gvisor.dev/gvisor/pkg/tcpip"
|
||||
"gvisor.dev/gvisor/pkg/tcpip/checksum"
|
||||
"gvisor.dev/gvisor/pkg/tcpip/header"
|
||||
"gvisor.dev/gvisor/pkg/tcpip/stack"
|
||||
)
|
||||
|
||||
// fakeGVisorDevice is an in-memory GVisorDevice used to script conversations
|
||||
// with the "system" stack without any real tun device or privileges.
|
||||
type fakeGVisorDevice struct {
|
||||
inbound chan []byte
|
||||
outbound chan []byte
|
||||
notify chan struct{}
|
||||
}
|
||||
|
||||
func newFakeGVisorDevice() *fakeGVisorDevice {
|
||||
return &fakeGVisorDevice{
|
||||
inbound: make(chan []byte, 256),
|
||||
outbound: make(chan []byte, 256),
|
||||
notify: make(chan struct{}, 1),
|
||||
}
|
||||
}
|
||||
|
||||
func (d *fakeGVisorDevice) push(data []byte) {
|
||||
d.inbound <- data
|
||||
select {
|
||||
case d.notify <- struct{}{}:
|
||||
default:
|
||||
}
|
||||
}
|
||||
|
||||
func (d *fakeGVisorDevice) ReadPacket() (byte, *stack.PacketBuffer, error) {
|
||||
select {
|
||||
case data := <-d.inbound:
|
||||
version := data[0] >> 4
|
||||
pkt := stack.NewPacketBuffer(stack.PacketBufferOptions{Payload: buffer.MakeWithData(data)})
|
||||
return version, pkt, nil
|
||||
default:
|
||||
return 0, nil, ErrQueueEmpty
|
||||
}
|
||||
}
|
||||
|
||||
func (d *fakeGVisorDevice) WritePacket(packet *stack.PacketBuffer) tcpip.Error {
|
||||
var data []byte
|
||||
for _, s := range packet.AsSlices() {
|
||||
data = append(data, s...)
|
||||
}
|
||||
d.outbound <- data
|
||||
return nil
|
||||
}
|
||||
|
||||
func (d *fakeGVisorDevice) Wait() {
|
||||
select {
|
||||
case <-d.notify:
|
||||
case <-time.After(20 * time.Millisecond):
|
||||
}
|
||||
}
|
||||
|
||||
func (d *fakeGVisorDevice) recv(t *testing.T, timeout time.Duration) []byte {
|
||||
t.Helper()
|
||||
select {
|
||||
case data := <-d.outbound:
|
||||
return data
|
||||
case <-time.After(timeout):
|
||||
t.Fatal("timed out waiting for outbound packet")
|
||||
return nil
|
||||
}
|
||||
}
|
||||
|
||||
var _ GVisorDevice = (*fakeGVisorDevice)(nil)
|
||||
|
||||
// echoHandler is a ConnectionHandler that echoes back everything it reads on
|
||||
// each connection, and records connections/destinations it has seen.
|
||||
type echoHandler struct {
|
||||
mu sync.Mutex
|
||||
conns []net.Conn
|
||||
dests []net.Destination
|
||||
done chan struct{}
|
||||
}
|
||||
|
||||
func newEchoHandler() *echoHandler {
|
||||
return &echoHandler{done: make(chan struct{}, 8)}
|
||||
}
|
||||
|
||||
func (h *echoHandler) HandleConnection(conn net.Conn, dest net.Destination) {
|
||||
h.mu.Lock()
|
||||
h.conns = append(h.conns, conn)
|
||||
h.dests = append(h.dests, dest)
|
||||
h.mu.Unlock()
|
||||
|
||||
_, _ = io.Copy(conn, conn)
|
||||
_ = conn.Close()
|
||||
h.done <- struct{}{}
|
||||
}
|
||||
|
||||
func newTestStackSystem(device GVisorDevice, handler ConnectionHandler, idleTimeout time.Duration) (*stackSystem, context.CancelFunc) {
|
||||
ctx, cancel := context.WithCancel(context.Background())
|
||||
s := &stackSystem{
|
||||
ctx: ctx,
|
||||
device: device,
|
||||
mtu: 1500,
|
||||
idleTimeout: idleTimeout,
|
||||
handler: handler,
|
||||
tcp: make(map[tcpKey]*tcpConn),
|
||||
}
|
||||
return s, cancel
|
||||
}
|
||||
|
||||
func testIP(s string) tcpip.Address {
|
||||
return tcpip.AddrFrom4Slice(net.ParseIP(s).To4())
|
||||
}
|
||||
|
||||
const (
|
||||
testPeerIP = "10.0.0.2"
|
||||
testTargetIP = "10.0.0.1"
|
||||
testPeerPort = uint16(51234)
|
||||
testDstPort = uint16(8080)
|
||||
)
|
||||
|
||||
// buildIPv4TCP builds a raw IPv4+TCP segment, computing valid checksums.
|
||||
func buildIPv4TCP(src, dst tcpip.Address, srcPort, dstPort uint16, seq, ack uint32, flags header.TCPFlags, window uint16, payload []byte, options []byte) []byte {
|
||||
headerLen := header.TCPMinimumSize + len(options)
|
||||
totalLen := header.IPv4MinimumSize + headerLen + len(payload)
|
||||
data := make([]byte, totalLen)
|
||||
|
||||
ipHdr := header.IPv4(data)
|
||||
ipHdr.Encode(&header.IPv4Fields{
|
||||
TotalLength: uint16(totalLen),
|
||||
TTL: 64,
|
||||
Protocol: uint8(header.TCPProtocolNumber),
|
||||
SrcAddr: src,
|
||||
DstAddr: dst,
|
||||
})
|
||||
ipHdr.SetChecksum(^ipHdr.CalculateChecksum())
|
||||
|
||||
tcpHdr := header.TCP(data[header.IPv4MinimumSize:])
|
||||
tcpHdr.Encode(&header.TCPFields{
|
||||
SrcPort: srcPort,
|
||||
DstPort: dstPort,
|
||||
SeqNum: seq,
|
||||
AckNum: ack,
|
||||
DataOffset: uint8(headerLen),
|
||||
Flags: flags,
|
||||
WindowSize: window,
|
||||
})
|
||||
copy(tcpHdr.Options(), options)
|
||||
copy(data[header.IPv4MinimumSize+headerLen:], payload)
|
||||
|
||||
xsum := header.PseudoHeaderChecksum(header.TCPProtocolNumber, src, dst, uint16(headerLen+len(payload)))
|
||||
xsum = checksum.Checksum(payload, xsum)
|
||||
tcpHdr.SetChecksum(^tcpHdr.CalculateChecksum(xsum))
|
||||
|
||||
return data
|
||||
}
|
||||
|
||||
func parseIPv4TCP(t *testing.T, data []byte) header.TCP {
|
||||
t.Helper()
|
||||
ipHdr := header.IPv4(data)
|
||||
if !ipHdr.IsValid(len(data)) {
|
||||
t.Fatalf("invalid ipv4 packet")
|
||||
}
|
||||
return header.TCP(ipHdr.Payload())
|
||||
}
|
||||
|
||||
func TestSystemStackTCPHandshakeEchoClose(t *testing.T) {
|
||||
device := newFakeGVisorDevice()
|
||||
handler := newEchoHandler()
|
||||
s, cancel := newTestStackSystem(device, handler, time.Minute)
|
||||
defer cancel()
|
||||
if err := s.Start(); err != nil {
|
||||
t.Fatalf("Start: %v", err)
|
||||
}
|
||||
defer s.Close()
|
||||
|
||||
src := testIP(testPeerIP)
|
||||
dst := testIP(testTargetIP)
|
||||
|
||||
iss := uint32(1000)
|
||||
device.push(buildIPv4TCP(src, dst, testPeerPort, testDstPort, iss, 0, header.TCPFlagSyn, 65535, nil, nil))
|
||||
|
||||
synAck := parseIPv4TCP(t, device.recv(t, time.Second))
|
||||
if synAck.Flags() != header.TCPFlagSyn|header.TCPFlagAck {
|
||||
t.Fatalf("expected SYN-ACK, got flags %v", synAck.Flags())
|
||||
}
|
||||
if synAck.AckNumber() != iss+1 {
|
||||
t.Fatalf("unexpected ack number %d, want %d", synAck.AckNumber(), iss+1)
|
||||
}
|
||||
serverISS := synAck.SequenceNumber()
|
||||
|
||||
// final handshake ACK
|
||||
device.push(buildIPv4TCP(src, dst, testPeerPort, testDstPort, iss+1, serverISS+1, header.TCPFlagAck, 65535, nil, nil))
|
||||
|
||||
// send data
|
||||
payload := []byte("hello world")
|
||||
device.push(buildIPv4TCP(src, dst, testPeerPort, testDstPort, iss+1, serverISS+1, header.TCPFlagAck|header.TCPFlagPsh, 65535, payload, nil))
|
||||
|
||||
// drain outbound packets until the full echo has been observed, acking
|
||||
// any data segments as they arrive so the connection can make progress
|
||||
var echoed []byte
|
||||
deadline := time.After(2 * time.Second)
|
||||
for len(echoed) < len(payload) {
|
||||
select {
|
||||
case raw := <-device.outbound:
|
||||
tcpHdr := parseIPv4TCP(t, raw)
|
||||
if len(tcpHdr.Payload()) > 0 {
|
||||
echoed = append(echoed, tcpHdr.Payload()...)
|
||||
device.push(buildIPv4TCP(src, dst, testPeerPort, testDstPort,
|
||||
iss+1+uint32(len(payload)), tcpHdr.SequenceNumber()+uint32(len(tcpHdr.Payload())),
|
||||
header.TCPFlagAck, 65535, nil, nil))
|
||||
}
|
||||
case <-deadline:
|
||||
t.Fatalf("timed out waiting for echo, got %q so far", echoed)
|
||||
}
|
||||
}
|
||||
if string(echoed) != string(payload) {
|
||||
t.Fatalf("echo mismatch: got %q want %q", echoed, payload)
|
||||
}
|
||||
|
||||
h := handler
|
||||
h.mu.Lock()
|
||||
if len(h.dests) != 1 || h.dests[0].NetAddr() != "10.0.0.1:8080" {
|
||||
t.Fatalf("unexpected destination recorded: %+v", h.dests)
|
||||
}
|
||||
h.mu.Unlock()
|
||||
|
||||
// peer sends FIN
|
||||
finSeq := iss + 1 + uint32(len(payload))
|
||||
device.push(buildIPv4TCP(src, dst, testPeerPort, testDstPort, finSeq, serverISS+1+uint32(len(payload)), header.TCPFlagFin|header.TCPFlagAck, 65535, nil, nil))
|
||||
|
||||
select {
|
||||
case <-handler.done:
|
||||
case <-time.After(2 * time.Second):
|
||||
t.Fatal("echo handler never finished after peer FIN")
|
||||
}
|
||||
|
||||
var sawAckOfFin, sawOurFin bool
|
||||
var ourFinSeq uint32
|
||||
deadline = time.After(2 * time.Second)
|
||||
for !sawAckOfFin || !sawOurFin {
|
||||
select {
|
||||
case raw := <-device.outbound:
|
||||
tcpHdr := parseIPv4TCP(t, raw)
|
||||
if tcpHdr.Flags()&header.TCPFlagFin != 0 {
|
||||
sawOurFin = true
|
||||
ourFinSeq = tcpHdr.SequenceNumber()
|
||||
}
|
||||
if tcpHdr.AckNumber() == finSeq+1 {
|
||||
sawAckOfFin = true
|
||||
}
|
||||
case <-deadline:
|
||||
t.Fatalf("timed out waiting for our fin/ack (sawAckOfFin=%v sawOurFin=%v)", sawAckOfFin, sawOurFin)
|
||||
}
|
||||
}
|
||||
|
||||
// ack our FIN, completing a graceful close
|
||||
device.push(buildIPv4TCP(src, dst, testPeerPort, testDstPort, finSeq+1, ourFinSeq+1, header.TCPFlagAck, 65535, nil, nil))
|
||||
|
||||
deadline = time.After(2 * time.Second)
|
||||
for {
|
||||
s.tcpMu.Lock()
|
||||
var conn *tcpConn
|
||||
for _, c := range s.tcp {
|
||||
conn = c
|
||||
}
|
||||
s.tcpMu.Unlock()
|
||||
if conn == nil {
|
||||
t.Fatal("connection unexpectedly removed before linger")
|
||||
}
|
||||
conn.mu.Lock()
|
||||
state := conn.state
|
||||
conn.mu.Unlock()
|
||||
if state == stateTimeWait {
|
||||
break
|
||||
}
|
||||
select {
|
||||
case <-deadline:
|
||||
t.Fatalf("connection did not reach TimeWait, state=%d", state)
|
||||
case <-time.After(10 * time.Millisecond):
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
func TestSystemStackTCPUnknownConnectionReset(t *testing.T) {
|
||||
device := newFakeGVisorDevice()
|
||||
handler := newEchoHandler()
|
||||
s, cancel := newTestStackSystem(device, handler, time.Minute)
|
||||
defer cancel()
|
||||
if err := s.Start(); err != nil {
|
||||
t.Fatalf("Start: %v", err)
|
||||
}
|
||||
defer s.Close()
|
||||
|
||||
src := testIP(testPeerIP)
|
||||
dst := testIP(testTargetIP)
|
||||
|
||||
// an ACK referencing a connection the stack has never seen
|
||||
device.push(buildIPv4TCP(src, dst, testPeerPort, testDstPort, 5000, 0, header.TCPFlagAck, 65535, nil, nil))
|
||||
|
||||
rst := parseIPv4TCP(t, device.recv(t, time.Second))
|
||||
if rst.Flags()&header.TCPFlagRst == 0 {
|
||||
t.Fatalf("expected RST, got flags %v", rst.Flags())
|
||||
}
|
||||
if rst.SequenceNumber() != 5000 {
|
||||
t.Fatalf("expected reset seq to echo the ack number 5000, got %d", rst.SequenceNumber())
|
||||
}
|
||||
}
|
||||
|
||||
func TestSystemStackTCPRetransmit(t *testing.T) {
|
||||
device := newFakeGVisorDevice()
|
||||
handler := newEchoHandler()
|
||||
s, cancel := newTestStackSystem(device, handler, time.Minute)
|
||||
defer cancel()
|
||||
if err := s.Start(); err != nil {
|
||||
t.Fatalf("Start: %v", err)
|
||||
}
|
||||
defer s.Close()
|
||||
|
||||
src := testIP(testPeerIP)
|
||||
dst := testIP(testTargetIP)
|
||||
|
||||
iss := uint32(2000)
|
||||
device.push(buildIPv4TCP(src, dst, testPeerPort, testDstPort, iss, 0, header.TCPFlagSyn, 65535, nil, nil))
|
||||
synAck := parseIPv4TCP(t, device.recv(t, time.Second))
|
||||
serverISS := synAck.SequenceNumber()
|
||||
device.push(buildIPv4TCP(src, dst, testPeerPort, testDstPort, iss+1, serverISS+1, header.TCPFlagAck, 65535, nil, nil))
|
||||
|
||||
payload := []byte("hi")
|
||||
device.push(buildIPv4TCP(src, dst, testPeerPort, testDstPort, iss+1, serverISS+1, header.TCPFlagAck|header.TCPFlagPsh, 65535, payload, nil))
|
||||
|
||||
// consume the data-ack and the first echoed data segment, but do NOT ack
|
||||
// the echoed data, forcing a retransmit
|
||||
var first []byte
|
||||
deadline := time.After(2 * time.Second)
|
||||
for len(first) == 0 {
|
||||
select {
|
||||
case raw := <-device.outbound:
|
||||
tcpHdr := parseIPv4TCP(t, raw)
|
||||
if len(tcpHdr.Payload()) > 0 {
|
||||
first = append([]byte(nil), tcpHdr.Payload()...)
|
||||
}
|
||||
case <-deadline:
|
||||
t.Fatal("timed out waiting for first echoed segment")
|
||||
}
|
||||
}
|
||||
|
||||
// now wait for a retransmission of the same bytes, without acking
|
||||
deadline = time.After(2 * time.Second)
|
||||
for {
|
||||
select {
|
||||
case raw := <-device.outbound:
|
||||
tcpHdr := parseIPv4TCP(t, raw)
|
||||
if string(tcpHdr.Payload()) == string(first) {
|
||||
return // retransmit observed, test passes
|
||||
}
|
||||
case <-deadline:
|
||||
t.Fatal("timed out waiting for retransmission of unacked data")
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
func TestSystemStackIdleReap(t *testing.T) {
|
||||
device := newFakeGVisorDevice()
|
||||
handler := newEchoHandler()
|
||||
s, cancel := newTestStackSystem(device, handler, time.Millisecond)
|
||||
defer cancel()
|
||||
if err := s.Start(); err != nil {
|
||||
t.Fatalf("Start: %v", err)
|
||||
}
|
||||
defer s.Close()
|
||||
|
||||
src := testIP(testPeerIP)
|
||||
dst := testIP(testTargetIP)
|
||||
|
||||
device.push(buildIPv4TCP(src, dst, testPeerPort, testDstPort, 1, 0, header.TCPFlagSyn, 65535, nil, nil))
|
||||
device.recv(t, time.Second) // SYN-ACK
|
||||
|
||||
deadline := time.After(2 * time.Second)
|
||||
for {
|
||||
s.tcpMu.Lock()
|
||||
n := len(s.tcp)
|
||||
s.tcpMu.Unlock()
|
||||
if n == 0 {
|
||||
return
|
||||
}
|
||||
select {
|
||||
case <-deadline:
|
||||
t.Fatal("idle connection was not reaped")
|
||||
case <-time.After(10 * time.Millisecond):
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
func TestSystemStackUDPEcho(t *testing.T) {
|
||||
device := newFakeGVisorDevice()
|
||||
handler := newEchoHandler()
|
||||
s, cancel := newTestStackSystem(device, handler, time.Minute)
|
||||
defer cancel()
|
||||
if err := s.Start(); err != nil {
|
||||
t.Fatalf("Start: %v", err)
|
||||
}
|
||||
defer s.Close()
|
||||
|
||||
src := testIP(testPeerIP)
|
||||
dst := testIP(testTargetIP)
|
||||
|
||||
payload := []byte("ping")
|
||||
udpLen := header.UDPMinimumSize + len(payload)
|
||||
totalLen := header.IPv4MinimumSize + udpLen
|
||||
data := make([]byte, totalLen)
|
||||
ipHdr := header.IPv4(data)
|
||||
ipHdr.Encode(&header.IPv4Fields{
|
||||
TotalLength: uint16(totalLen),
|
||||
TTL: 64,
|
||||
Protocol: uint8(header.UDPProtocolNumber),
|
||||
SrcAddr: src,
|
||||
DstAddr: dst,
|
||||
})
|
||||
ipHdr.SetChecksum(^ipHdr.CalculateChecksum())
|
||||
udpHdr := header.UDP(data[header.IPv4MinimumSize:])
|
||||
udpHdr.Encode(&header.UDPFields{SrcPort: testPeerPort, DstPort: testDstPort, Length: uint16(udpLen)})
|
||||
copy(data[header.IPv4MinimumSize+header.UDPMinimumSize:], payload)
|
||||
xsum := header.PseudoHeaderChecksum(header.UDPProtocolNumber, src, dst, uint16(udpLen))
|
||||
udpHdr.SetChecksum(^udpHdr.CalculateChecksum(checksum.Checksum(payload, xsum)))
|
||||
|
||||
device.push(data)
|
||||
|
||||
select {
|
||||
case <-handler.done:
|
||||
case <-time.After(time.Second):
|
||||
t.Fatal("udp handler never invoked/finished")
|
||||
}
|
||||
|
||||
handler.mu.Lock()
|
||||
defer handler.mu.Unlock()
|
||||
if len(handler.dests) != 1 || handler.dests[0].Network != net.Network_UDP {
|
||||
t.Fatalf("unexpected udp destination recorded: %+v", handler.dests)
|
||||
}
|
||||
}
|
||||
@@ -7,9 +7,12 @@ import (
|
||||
"net"
|
||||
"strconv"
|
||||
|
||||
"github.com/xtls/xray-core/common/buf"
|
||||
"github.com/xtls/xray-core/common/errors"
|
||||
"github.com/xtls/xray-core/common/platform"
|
||||
"golang.org/x/sys/unix"
|
||||
"gvisor.dev/gvisor/pkg/buffer"
|
||||
"gvisor.dev/gvisor/pkg/tcpip"
|
||||
"gvisor.dev/gvisor/pkg/tcpip/link/fdbased"
|
||||
"gvisor.dev/gvisor/pkg/tcpip/stack"
|
||||
)
|
||||
@@ -22,6 +25,22 @@ type AndroidTun struct {
|
||||
// DefaultTun implements Tun
|
||||
var _ Tun = (*AndroidTun)(nil)
|
||||
|
||||
// AndroidTun implements GVisorDevice, used by the "system" (lite) ip stack
|
||||
var _ GVisorDevice = (*AndroidTun)(nil)
|
||||
|
||||
// fdReadWriter adapts a raw, already non-blocking file descriptor to io.Reader/io.Writer,
|
||||
// so it can be used with buf.Buffer.ReadFrom, without the ownership/finalizer overhead of
|
||||
// wrapping it in an *os.File (the fd is owned and closed elsewhere).
|
||||
type fdReadWriter int
|
||||
|
||||
func (f fdReadWriter) Read(p []byte) (int, error) {
|
||||
return unix.Read(int(f), p)
|
||||
}
|
||||
|
||||
func (f fdReadWriter) Write(p []byte) (int, error) {
|
||||
return unix.Write(int(f), p)
|
||||
}
|
||||
|
||||
// NewTun builds new tun interface handler
|
||||
func NewTun(options *Config) (Tun, error) {
|
||||
fd, err := strconv.Atoi(platform.NewEnvFlag(platform.TunFdKey).GetValue(func() string { return "0" }))
|
||||
@@ -78,6 +97,72 @@ func (t *AndroidTun) newEndpoint() (stack.LinkEndpoint, error) {
|
||||
})
|
||||
}
|
||||
|
||||
// ReadPacket implements GVisorDevice method to read one packet from the tun device, used by
|
||||
// the "system" (lite) ip stack. The gVisor backed stack instead talks to the fd directly through
|
||||
// fdbased.New above, for lower overhead batched IO, bypassing GVisorDevice entirely.
|
||||
// It is expected that the method will not block, rather return ErrQueueEmpty when there is nothing on the line,
|
||||
// which will make the stack call Wait which should implement desired push-back
|
||||
func (t *AndroidTun) ReadPacket() (byte, *stack.PacketBuffer, error) {
|
||||
// request memory to write from reusable buffer pool
|
||||
b := buf.NewWithSize(int32(t.options.MTU))
|
||||
|
||||
// read the bytes from the interface file descriptor, which is already non-blocking
|
||||
n, err := b.ReadFrom(fdReadWriter(t.tunFd))
|
||||
if err == unix.EAGAIN || err == unix.EWOULDBLOCK || err == unix.EINTR {
|
||||
b.Release()
|
||||
return 0, nil, ErrQueueEmpty
|
||||
}
|
||||
if err != nil {
|
||||
b.Release()
|
||||
return 0, nil, err
|
||||
}
|
||||
|
||||
// discard empty packets
|
||||
if n == 0 {
|
||||
b.Release()
|
||||
return 0, nil, ErrQueueEmpty
|
||||
}
|
||||
|
||||
// network protocol version from the first nibble of the raw packet
|
||||
version := b.Byte(0) >> 4
|
||||
packetBuffer := buffer.MakeWithData(b.Bytes())
|
||||
return version, stack.NewPacketBuffer(stack.PacketBufferOptions{
|
||||
Payload: packetBuffer,
|
||||
IsForwardedPacket: true,
|
||||
OnRelease: func() {
|
||||
b.Release()
|
||||
},
|
||||
}), nil
|
||||
}
|
||||
|
||||
// WritePacket implements GVisorDevice method to write one packet to the tun device
|
||||
func (t *AndroidTun) WritePacket(packet *stack.PacketBuffer) tcpip.Error {
|
||||
// request memory to write from reusable buffer pool
|
||||
b := buf.NewWithSize(int32(t.options.MTU))
|
||||
defer b.Release()
|
||||
|
||||
// copy the bytes of slices that compose the packet into the allocated buffer, no
|
||||
// extra header is needed here, unlike Darwin/FreeBSD's utun devices
|
||||
for _, packetElement := range packet.AsSlices() {
|
||||
_, _ = b.Write(packetElement)
|
||||
}
|
||||
|
||||
if _, err := fdReadWriter(t.tunFd).Write(b.Bytes()); err != nil {
|
||||
if err == unix.EAGAIN || err == unix.EWOULDBLOCK {
|
||||
return &tcpip.ErrWouldBlock{}
|
||||
}
|
||||
return &tcpip.ErrAborted{}
|
||||
}
|
||||
return nil
|
||||
}
|
||||
|
||||
// Wait blocks until the tun fd is likely readable again, rather than spinning the CPU.
|
||||
// A bounded timeout keeps this responsive to a Close() racing a call already parked here.
|
||||
func (t *AndroidTun) Wait() {
|
||||
fds := []unix.PollFd{{Fd: int32(t.tunFd), Events: unix.POLLIN}}
|
||||
_, _ = unix.Poll(fds, 1000)
|
||||
}
|
||||
|
||||
func setinterface(network, address string, fd uintptr, iface *net.Interface) error {
|
||||
return unix.BindToDevice(int(fd), iface.Name)
|
||||
}
|
||||
|
||||
@@ -10,9 +10,12 @@ import (
|
||||
"sync"
|
||||
|
||||
"github.com/vishvananda/netlink"
|
||||
"github.com/xtls/xray-core/common/buf"
|
||||
"github.com/xtls/xray-core/common/errors"
|
||||
"github.com/xtls/xray-core/common/platform"
|
||||
"golang.org/x/sys/unix"
|
||||
"gvisor.dev/gvisor/pkg/buffer"
|
||||
"gvisor.dev/gvisor/pkg/tcpip"
|
||||
"gvisor.dev/gvisor/pkg/tcpip/link/fdbased"
|
||||
"gvisor.dev/gvisor/pkg/tcpip/stack"
|
||||
)
|
||||
@@ -35,6 +38,22 @@ type LinuxTun struct {
|
||||
// LinuxTun implements Tun
|
||||
var _ Tun = (*LinuxTun)(nil)
|
||||
|
||||
// LinuxTun implements GVisorDevice, used by the "system" (lite) ip stack
|
||||
var _ GVisorDevice = (*LinuxTun)(nil)
|
||||
|
||||
// fdReadWriter adapts a raw, already non-blocking file descriptor to io.Reader/io.Writer,
|
||||
// so it can be used with buf.Buffer.ReadFrom, without the ownership/finalizer overhead of
|
||||
// wrapping it in an *os.File (the fd is owned and closed elsewhere, see LinuxTun.Close).
|
||||
type fdReadWriter int
|
||||
|
||||
func (f fdReadWriter) Read(p []byte) (int, error) {
|
||||
return unix.Read(int(f), p)
|
||||
}
|
||||
|
||||
func (f fdReadWriter) Write(p []byte) (int, error) {
|
||||
return unix.Write(int(f), p)
|
||||
}
|
||||
|
||||
// NewTun builds new tun interface handler (linux specific)
|
||||
func NewTun(options *Config) (Tun, error) {
|
||||
tunFd, tunLink, fdProvided, err := openFromEnv(options.Name)
|
||||
@@ -228,6 +247,72 @@ func (t *LinuxTun) newEndpoint() (stack.LinkEndpoint, error) {
|
||||
})
|
||||
}
|
||||
|
||||
// ReadPacket implements GVisorDevice method to read one packet from the tun device, used by
|
||||
// the "system" (lite) ip stack. The gVisor backed stack instead talks to the fd directly through
|
||||
// fdbased.New above, for lower overhead batched IO, bypassing GVisorDevice entirely.
|
||||
// It is expected that the method will not block, rather return ErrQueueEmpty when there is nothing on the line,
|
||||
// which will make the stack call Wait which should implement desired push-back
|
||||
func (t *LinuxTun) ReadPacket() (byte, *stack.PacketBuffer, error) {
|
||||
// request memory to write from reusable buffer pool
|
||||
b := buf.NewWithSize(int32(t.options.MTU))
|
||||
|
||||
// read the bytes from the interface file descriptor, which is already non-blocking
|
||||
n, err := b.ReadFrom(fdReadWriter(t.tunFd))
|
||||
if err == unix.EAGAIN || err == unix.EWOULDBLOCK || err == unix.EINTR {
|
||||
b.Release()
|
||||
return 0, nil, ErrQueueEmpty
|
||||
}
|
||||
if err != nil {
|
||||
b.Release()
|
||||
return 0, nil, err
|
||||
}
|
||||
|
||||
// discard empty packets
|
||||
if n == 0 {
|
||||
b.Release()
|
||||
return 0, nil, ErrQueueEmpty
|
||||
}
|
||||
|
||||
// network protocol version from the first nibble of the raw packet
|
||||
version := b.Byte(0) >> 4
|
||||
packetBuffer := buffer.MakeWithData(b.Bytes())
|
||||
return version, stack.NewPacketBuffer(stack.PacketBufferOptions{
|
||||
Payload: packetBuffer,
|
||||
IsForwardedPacket: true,
|
||||
OnRelease: func() {
|
||||
b.Release()
|
||||
},
|
||||
}), nil
|
||||
}
|
||||
|
||||
// WritePacket implements GVisorDevice method to write one packet to the tun device
|
||||
func (t *LinuxTun) WritePacket(packet *stack.PacketBuffer) tcpip.Error {
|
||||
// request memory to write from reusable buffer pool
|
||||
b := buf.NewWithSize(int32(t.options.MTU))
|
||||
defer b.Release()
|
||||
|
||||
// copy the bytes of slices that compose the packet into the allocated buffer, no
|
||||
// Linux specific header is needed here, unlike Darwin/FreeBSD's utun devices
|
||||
for _, packetElement := range packet.AsSlices() {
|
||||
_, _ = b.Write(packetElement)
|
||||
}
|
||||
|
||||
if _, err := fdReadWriter(t.tunFd).Write(b.Bytes()); err != nil {
|
||||
if err == unix.EAGAIN || err == unix.EWOULDBLOCK {
|
||||
return &tcpip.ErrWouldBlock{}
|
||||
}
|
||||
return &tcpip.ErrAborted{}
|
||||
}
|
||||
return nil
|
||||
}
|
||||
|
||||
// Wait blocks until the tun fd is likely readable again, rather than spinning the CPU.
|
||||
// A bounded timeout keeps this responsive to a Close() racing a call already parked here.
|
||||
func (t *LinuxTun) Wait() {
|
||||
fds := []unix.PollFd{{Fd: int32(t.tunFd), Events: unix.POLLIN}}
|
||||
_, _ = unix.Poll(fds, 1000)
|
||||
}
|
||||
|
||||
func setinterface(network, address string, fd uintptr, iface *net.Interface) error {
|
||||
return unix.BindToDevice(int(fd), iface.Name)
|
||||
}
|
||||
|
||||
+111
-136
@@ -3,7 +3,6 @@ package wireguard
|
||||
import (
|
||||
"context"
|
||||
"fmt"
|
||||
gonet "net"
|
||||
"net/netip"
|
||||
"reflect"
|
||||
"strings"
|
||||
@@ -28,14 +27,10 @@ import (
|
||||
"github.com/xtls/xray-core/features/stats"
|
||||
"github.com/xtls/xray-core/transport"
|
||||
"github.com/xtls/xray-core/transport/internet"
|
||||
"github.com/xtls/xray-core/transport/internet/finalmask"
|
||||
"golang.zx2c4.com/wireguard/device"
|
||||
)
|
||||
|
||||
type entry struct {
|
||||
got []net.IP
|
||||
time time.Time
|
||||
}
|
||||
|
||||
type Handler struct {
|
||||
conf *DeviceConfig
|
||||
policyManager policy.Manager
|
||||
@@ -49,11 +44,6 @@ type Handler struct {
|
||||
tnet *Net
|
||||
dev *device.Device
|
||||
mu sync.Mutex
|
||||
|
||||
// TODO: cache cleanup loop
|
||||
local bool
|
||||
cache map[string]entry
|
||||
cacheMu sync.Mutex
|
||||
}
|
||||
|
||||
func NewClient(ctx context.Context, conf *DeviceConfig) (*Handler, error) {
|
||||
@@ -109,15 +99,10 @@ func NewClient(ctx context.Context, conf *DeviceConfig) (*Handler, error) {
|
||||
return nil, err
|
||||
}
|
||||
|
||||
local := false
|
||||
dns := conf.DNS
|
||||
if len(dns) == 0 {
|
||||
dns = []string{"1.1.1.1", "1.0.0.1", "2606:4700:4700::1111", "2606:4700:4700::1001"}
|
||||
}
|
||||
if len(dns) == 1 && dns[0] == "local" {
|
||||
local = true
|
||||
dns = nil
|
||||
}
|
||||
dnses := make([]netip.Addr, 0, len(dns))
|
||||
for _, dns := range dns {
|
||||
dnses = append(dnses, netip.MustParseAddr(dns))
|
||||
@@ -151,9 +136,6 @@ func NewClient(ctx context.Context, conf *DeviceConfig) (*Handler, error) {
|
||||
|
||||
tun: tun,
|
||||
tnet: tnet,
|
||||
|
||||
local: local,
|
||||
cache: make(map[string]entry),
|
||||
}, nil
|
||||
}
|
||||
|
||||
@@ -172,22 +154,6 @@ func (h *Handler) Process(ctx context.Context, link *transport.Link, dialer inte
|
||||
return err
|
||||
}
|
||||
|
||||
var addr netip.Addr
|
||||
if ob.Target.Address.Family().IsDomain() {
|
||||
ip, err := h.resolveRemote(ob.Target.Address.String())
|
||||
if err != nil {
|
||||
return errors.New("failed to resolve domain").Base(err)
|
||||
}
|
||||
addr, _ = netip.AddrFromSlice(ip)
|
||||
} else {
|
||||
addr, _ = netip.AddrFromSlice(ob.Target.Address.IP())
|
||||
}
|
||||
|
||||
addrPort := netip.AddrPortFrom(addr, ob.Target.Port.Value())
|
||||
if !addrPort.IsValid() {
|
||||
return errors.New("invalid target ", ob.Target)
|
||||
}
|
||||
|
||||
var newCtx context.Context
|
||||
var newCancel context.CancelFunc
|
||||
if session.TimeoutOnlyFromContext(ctx) {
|
||||
@@ -216,10 +182,10 @@ func (h *Handler) Process(ctx context.Context, link *transport.Link, dialer inte
|
||||
var err error
|
||||
if sessionPolicy.Timeouts.Handshake != 0 {
|
||||
timeoutCtx, timeoutCancel := context.WithTimeout(ctx, sessionPolicy.Timeouts.Handshake)
|
||||
conn, err = h.tnet.DialContextTCPAddrPort(timeoutCtx, addrPort)
|
||||
conn, err = h.tnet.DialContext(timeoutCtx, "tcp", ob.Target.NetAddr())
|
||||
timeoutCancel()
|
||||
} else {
|
||||
conn, err = h.tnet.DialContextTCPAddrPort(ctx, addrPort)
|
||||
conn, err = h.tnet.Dial("tcp", ob.Target.NetAddr())
|
||||
}
|
||||
if err != nil {
|
||||
return errors.New("failed to create TCP connection").Base(err)
|
||||
@@ -228,15 +194,14 @@ func (h *Handler) Process(ctx context.Context, link *transport.Link, dialer inte
|
||||
reader = buf.NewReader(conn)
|
||||
writer = buf.NewWriter(conn)
|
||||
case net.Network_UDP:
|
||||
conn, err := h.tnet.DialUDPAddrPort(netip.AddrPort{}, addrPort)
|
||||
conn, err := h.tnet.Dial("udp", ob.Target.NetAddr())
|
||||
if err != nil {
|
||||
return errors.New("failed to create UDP connection").Base(err)
|
||||
}
|
||||
defer conn.Close()
|
||||
c := &udpConnClient{
|
||||
PacketConn: conn.(*internet.PacketConnWrapper).PacketConn,
|
||||
resolveFunc: h.resolveRemote,
|
||||
dest: gonet.UDPAddrFromAddrPort(addrPort),
|
||||
PacketConn: conn.(*internet.PacketConnWrapper).PacketConn,
|
||||
dest: conn.RemoteAddr().(*net.UDPAddr),
|
||||
}
|
||||
reader = c
|
||||
writer = c
|
||||
@@ -293,26 +258,26 @@ func (h *Handler) init(ctx context.Context) error {
|
||||
if err != nil {
|
||||
return nil, err
|
||||
}
|
||||
conn, err := internet.DialSystem(ctx, dest, h.streamSettings.SocketSettings)
|
||||
if err != nil {
|
||||
return nil, err
|
||||
}
|
||||
var pktConn net.PacketConn
|
||||
switch c := conn.(type) {
|
||||
case *internet.PacketConnWrapper:
|
||||
pktConn = c.PacketConn
|
||||
case *cnc.Connection:
|
||||
pktConn = &internet.FakePacketConn{Conn: c}
|
||||
default:
|
||||
panic(reflect.TypeOf(c))
|
||||
}
|
||||
if h.streamSettings.UdpmaskManager != nil {
|
||||
newConn, err := h.streamSettings.UdpmaskManager.WrapPacketConnClient(pktConn)
|
||||
if h.streamSettings.FinalMask != nil {
|
||||
conn, err := h.streamSettings.FinalMask.DialUDP(ctx, dest)
|
||||
if err != nil {
|
||||
pktConn.Close()
|
||||
return nil, errors.New("mask err").Base(err)
|
||||
return nil, errors.New("failed to dial to dest").Base(err)
|
||||
}
|
||||
pktConn = conn.(*finalmask.PacketConnWrapper).PacketConn
|
||||
} else {
|
||||
conn, err := internet.DialSystem(ctx, dest, h.streamSettings.SocketSettings)
|
||||
if err != nil {
|
||||
return nil, errors.New("failed to dial to dest").Base(err)
|
||||
}
|
||||
switch c := conn.(type) {
|
||||
case *internet.PacketConnWrapper:
|
||||
pktConn = c.PacketConn
|
||||
case *cnc.Connection:
|
||||
pktConn = &internet.FakePacketConn{Conn: c}
|
||||
default:
|
||||
panic(reflect.TypeOf(c))
|
||||
}
|
||||
pktConn = newConn
|
||||
}
|
||||
if h.uplinkCounter != nil || h.downlinkCounter != nil {
|
||||
pktConn = &PacketCounterConnection{
|
||||
@@ -371,87 +336,48 @@ func (h *Handler) init(ctx context.Context) error {
|
||||
}
|
||||
|
||||
func (h *Handler) resolveLocal(host string) (net.IP, error) {
|
||||
return h.resolveDomain(host, h.conf.DomainStrategy, func(host string) ([]net.IP, uint32, error) {
|
||||
return h.dns.LookupIP(host, dns.IPOption{IPv4Enable: true, IPv6Enable: true})
|
||||
})
|
||||
}
|
||||
|
||||
func (h *Handler) resolveRemote(host string) (net.IP, error) {
|
||||
return h.resolveDomain(host, h.conf.DomainStrategy, func(host string) ([]net.IP, uint32, error) {
|
||||
if h.local {
|
||||
return h.dns.LookupIP(host, dns.IPOption{IPv4Enable: true, IPv6Enable: true})
|
||||
}
|
||||
return h.tnet.LookupHost(host)
|
||||
})
|
||||
}
|
||||
|
||||
func (h *Handler) resolveDomain(host string, strategy DeviceConfig_DomainStrategy, lookupIP func(host string) ([]net.IP, uint32, error)) (net.IP, error) {
|
||||
if ip := net.ParseIP(host); ip != nil {
|
||||
return ip, nil
|
||||
}
|
||||
h.cacheMu.Lock()
|
||||
if entry, ok := h.cache[host]; ok {
|
||||
if time.Now().Before(entry.time) {
|
||||
h.cacheMu.Unlock()
|
||||
return entry.got[dice.Roll(len(entry.got))], nil
|
||||
}
|
||||
delete(h.cache, host)
|
||||
}
|
||||
h.cacheMu.Unlock()
|
||||
ips, ttl, err := lookupIP(host)
|
||||
ips, _, err := h.dns.LookupIP(host, dns.IPOption{IPv4Enable: true, IPv6Enable: true})
|
||||
if err != nil {
|
||||
return nil, err
|
||||
}
|
||||
if len(ips) == 0 {
|
||||
return nil, dns.ErrEmptyResponse
|
||||
}
|
||||
var got4, got6 []net.IP
|
||||
for _, ip := range ips {
|
||||
if ip.To4() != nil {
|
||||
got4 = append(got4, ip)
|
||||
} else {
|
||||
got6 = append(got6, ip)
|
||||
got := ips
|
||||
if h.streamSettings.SocketSettings != nil {
|
||||
var got4, got6 []net.IP
|
||||
for _, ip := range ips {
|
||||
if ip.To4() != nil {
|
||||
got4 = append(got4, ip)
|
||||
} else {
|
||||
got6 = append(got6, ip)
|
||||
}
|
||||
}
|
||||
}
|
||||
var got []net.IP
|
||||
switch strategy {
|
||||
case DeviceConfig_FORCE_IP:
|
||||
got = ips
|
||||
return ips[dice.Roll(len(ips))], nil
|
||||
case DeviceConfig_FORCE_IP4:
|
||||
got = got4
|
||||
case DeviceConfig_FORCE_IP6:
|
||||
got = got6
|
||||
case DeviceConfig_FORCE_IP46:
|
||||
got = got4
|
||||
if len(got) == 0 {
|
||||
got = got6
|
||||
}
|
||||
case DeviceConfig_FORCE_IP64:
|
||||
got = got6
|
||||
if len(got) == 0 {
|
||||
switch h.streamSettings.SocketSettings.DomainStrategy {
|
||||
case internet.DomainStrategy_AS_IS, internet.DomainStrategy_USE_IP, internet.DomainStrategy_FORCE_IP:
|
||||
got = ips
|
||||
case internet.DomainStrategy_USE_IP4, internet.DomainStrategy_FORCE_IP4:
|
||||
got = got4
|
||||
case internet.DomainStrategy_USE_IP6, internet.DomainStrategy_FORCE_IP6:
|
||||
got = got6
|
||||
case internet.DomainStrategy_USE_IP46, internet.DomainStrategy_FORCE_IP46:
|
||||
got = got4
|
||||
if len(got) == 0 {
|
||||
got = got6
|
||||
}
|
||||
case internet.DomainStrategy_USE_IP64, internet.DomainStrategy_FORCE_IP64:
|
||||
got = got6
|
||||
if len(got) == 0 {
|
||||
got = got4
|
||||
}
|
||||
}
|
||||
if len(got) == 0 {
|
||||
return nil, dns.ErrEmptyResponse
|
||||
}
|
||||
default:
|
||||
panic(strategy)
|
||||
}
|
||||
if len(got) == 0 {
|
||||
return nil, dns.ErrEmptyResponse
|
||||
}
|
||||
entry := entry{
|
||||
got: got,
|
||||
time: time.Now().Add(time.Duration(ttl) * time.Second),
|
||||
}
|
||||
h.cacheMu.Lock()
|
||||
h.cache[host] = entry
|
||||
h.cacheMu.Unlock()
|
||||
return got[dice.Roll(len(got))], nil
|
||||
}
|
||||
|
||||
type udpConnClient struct {
|
||||
net.PacketConn
|
||||
resolveFunc func(host string) (net.IP, error)
|
||||
dest *net.UDPAddr
|
||||
dest *net.UDPAddr
|
||||
}
|
||||
|
||||
func (c *udpConnClient) ReadMultiBuffer() (buf.MultiBuffer, error) {
|
||||
@@ -478,15 +404,8 @@ func (c *udpConnClient) WriteMultiBuffer(mb buf.MultiBuffer) error {
|
||||
dst := c.dest
|
||||
if b.UDP != nil {
|
||||
if b.UDP.Address.Family().IsDomain() {
|
||||
ip, err := c.resolveFunc(b.UDP.Address.String())
|
||||
if err != nil {
|
||||
errors.LogErrorInner(context.Background(), err, "drop packet to ", b.UDP, " with size ", len(b.Bytes()))
|
||||
b.Release()
|
||||
continue
|
||||
}
|
||||
dst = &net.UDPAddr{
|
||||
IP: ip,
|
||||
Port: int(b.UDP.Port),
|
||||
if b.UDP.Port != net.Port(dst.Port) {
|
||||
dst = &net.UDPAddr{IP: dst.IP, Port: int(b.UDP.Port)}
|
||||
}
|
||||
} else {
|
||||
dst = b.UDP.RawNetAddr().(*net.UDPAddr)
|
||||
@@ -523,3 +442,59 @@ func (c *PacketCounterConnection) WriteTo(p []byte, addr net.Addr) (n int, err e
|
||||
}
|
||||
return
|
||||
}
|
||||
|
||||
type entry struct {
|
||||
saddr []string
|
||||
deadline time.Time
|
||||
}
|
||||
|
||||
type cache struct {
|
||||
running bool
|
||||
m map[string]entry
|
||||
mu sync.Mutex
|
||||
}
|
||||
|
||||
func (c *cache) run() {
|
||||
if c.running {
|
||||
return
|
||||
}
|
||||
c.running = true
|
||||
c.m = make(map[string]entry)
|
||||
go c.gc()
|
||||
}
|
||||
|
||||
func (c *cache) gc() {
|
||||
ticker := time.NewTicker(time.Minute)
|
||||
for {
|
||||
now := <-ticker.C
|
||||
c.mu.Lock()
|
||||
for key, entry := range c.m {
|
||||
if now.After(entry.deadline) {
|
||||
delete(c.m, key)
|
||||
}
|
||||
}
|
||||
c.mu.Unlock()
|
||||
}
|
||||
}
|
||||
|
||||
func (c *cache) LookupHost(host string) []string {
|
||||
c.mu.Lock()
|
||||
defer c.mu.Unlock()
|
||||
c.run()
|
||||
if entry, ok := c.m[host]; ok {
|
||||
if time.Now().Before(entry.deadline) {
|
||||
return entry.saddr
|
||||
}
|
||||
delete(c.m, host)
|
||||
}
|
||||
return nil
|
||||
}
|
||||
|
||||
func (c *cache) Cache(host string, saddr []string, ttl uint32) {
|
||||
c.mu.Lock()
|
||||
defer c.mu.Unlock()
|
||||
c.m[host] = entry{
|
||||
saddr: saddr,
|
||||
deadline: time.Now().Add(time.Second * time.Duration(ttl)),
|
||||
}
|
||||
}
|
||||
|
||||
+26
-102
@@ -22,61 +22,6 @@ const (
|
||||
_ = protoimpl.EnforceVersion(protoimpl.MaxVersion - 20)
|
||||
)
|
||||
|
||||
type DeviceConfig_DomainStrategy int32
|
||||
|
||||
const (
|
||||
DeviceConfig_FORCE_IP DeviceConfig_DomainStrategy = 0
|
||||
DeviceConfig_FORCE_IP4 DeviceConfig_DomainStrategy = 1
|
||||
DeviceConfig_FORCE_IP6 DeviceConfig_DomainStrategy = 2
|
||||
DeviceConfig_FORCE_IP46 DeviceConfig_DomainStrategy = 3
|
||||
DeviceConfig_FORCE_IP64 DeviceConfig_DomainStrategy = 4
|
||||
)
|
||||
|
||||
// Enum value maps for DeviceConfig_DomainStrategy.
|
||||
var (
|
||||
DeviceConfig_DomainStrategy_name = map[int32]string{
|
||||
0: "FORCE_IP",
|
||||
1: "FORCE_IP4",
|
||||
2: "FORCE_IP6",
|
||||
3: "FORCE_IP46",
|
||||
4: "FORCE_IP64",
|
||||
}
|
||||
DeviceConfig_DomainStrategy_value = map[string]int32{
|
||||
"FORCE_IP": 0,
|
||||
"FORCE_IP4": 1,
|
||||
"FORCE_IP6": 2,
|
||||
"FORCE_IP46": 3,
|
||||
"FORCE_IP64": 4,
|
||||
}
|
||||
)
|
||||
|
||||
func (x DeviceConfig_DomainStrategy) Enum() *DeviceConfig_DomainStrategy {
|
||||
p := new(DeviceConfig_DomainStrategy)
|
||||
*p = x
|
||||
return p
|
||||
}
|
||||
|
||||
func (x DeviceConfig_DomainStrategy) String() string {
|
||||
return protoimpl.X.EnumStringOf(x.Descriptor(), protoreflect.EnumNumber(x))
|
||||
}
|
||||
|
||||
func (DeviceConfig_DomainStrategy) Descriptor() protoreflect.EnumDescriptor {
|
||||
return file_proxy_wireguard_config_proto_enumTypes[0].Descriptor()
|
||||
}
|
||||
|
||||
func (DeviceConfig_DomainStrategy) Type() protoreflect.EnumType {
|
||||
return &file_proxy_wireguard_config_proto_enumTypes[0]
|
||||
}
|
||||
|
||||
func (x DeviceConfig_DomainStrategy) Number() protoreflect.EnumNumber {
|
||||
return protoreflect.EnumNumber(x)
|
||||
}
|
||||
|
||||
// Deprecated: Use DeviceConfig_DomainStrategy.Descriptor instead.
|
||||
func (DeviceConfig_DomainStrategy) EnumDescriptor() ([]byte, []int) {
|
||||
return file_proxy_wireguard_config_proto_rawDescGZIP(), []int{1, 0}
|
||||
}
|
||||
|
||||
type PeerConfig struct {
|
||||
state protoimpl.MessageState `protogen:"open.v1"`
|
||||
PublicKey string `protobuf:"bytes,1,opt,name=public_key,json=publicKey,proto3" json:"public_key,omitempty"`
|
||||
@@ -154,19 +99,18 @@ func (x *PeerConfig) GetAllowedIps() []string {
|
||||
}
|
||||
|
||||
type DeviceConfig struct {
|
||||
state protoimpl.MessageState `protogen:"open.v1"`
|
||||
SecretKey string `protobuf:"bytes,1,opt,name=secret_key,json=secretKey,proto3" json:"secret_key,omitempty"`
|
||||
Endpoint []string `protobuf:"bytes,2,rep,name=endpoint,proto3" json:"endpoint,omitempty"`
|
||||
Peers []*PeerConfig `protobuf:"bytes,3,rep,name=peers,proto3" json:"peers,omitempty"`
|
||||
Users []*protocol.User `protobuf:"bytes,5,rep,name=users,proto3" json:"users,omitempty"`
|
||||
Mtu int32 `protobuf:"varint,4,opt,name=mtu,proto3" json:"mtu,omitempty"`
|
||||
Reserved []byte `protobuf:"bytes,6,opt,name=reserved,proto3" json:"reserved,omitempty"`
|
||||
DomainStrategy DeviceConfig_DomainStrategy `protobuf:"varint,7,opt,name=domain_strategy,json=domainStrategy,proto3,enum=xray.proxy.wireguard.DeviceConfig_DomainStrategy" json:"domain_strategy,omitempty"`
|
||||
IsClient bool `protobuf:"varint,8,opt,name=is_client,json=isClient,proto3" json:"is_client,omitempty"`
|
||||
NoKernelTun bool `protobuf:"varint,9,opt,name=no_kernel_tun,json=noKernelTun,proto3" json:"no_kernel_tun,omitempty"`
|
||||
DNS []string `protobuf:"bytes,10,rep,name=DNS,proto3" json:"DNS,omitempty"`
|
||||
unknownFields protoimpl.UnknownFields
|
||||
sizeCache protoimpl.SizeCache
|
||||
state protoimpl.MessageState `protogen:"open.v1"`
|
||||
SecretKey string `protobuf:"bytes,1,opt,name=secret_key,json=secretKey,proto3" json:"secret_key,omitempty"`
|
||||
Endpoint []string `protobuf:"bytes,2,rep,name=endpoint,proto3" json:"endpoint,omitempty"`
|
||||
Peers []*PeerConfig `protobuf:"bytes,3,rep,name=peers,proto3" json:"peers,omitempty"`
|
||||
Users []*protocol.User `protobuf:"bytes,5,rep,name=users,proto3" json:"users,omitempty"`
|
||||
Mtu int32 `protobuf:"varint,4,opt,name=mtu,proto3" json:"mtu,omitempty"`
|
||||
Reserved []byte `protobuf:"bytes,6,opt,name=reserved,proto3" json:"reserved,omitempty"`
|
||||
IsClient bool `protobuf:"varint,8,opt,name=is_client,json=isClient,proto3" json:"is_client,omitempty"`
|
||||
NoKernelTun bool `protobuf:"varint,9,opt,name=no_kernel_tun,json=noKernelTun,proto3" json:"no_kernel_tun,omitempty"`
|
||||
DNS []string `protobuf:"bytes,10,rep,name=DNS,proto3" json:"DNS,omitempty"`
|
||||
unknownFields protoimpl.UnknownFields
|
||||
sizeCache protoimpl.SizeCache
|
||||
}
|
||||
|
||||
func (x *DeviceConfig) Reset() {
|
||||
@@ -241,13 +185,6 @@ func (x *DeviceConfig) GetReserved() []byte {
|
||||
return nil
|
||||
}
|
||||
|
||||
func (x *DeviceConfig) GetDomainStrategy() DeviceConfig_DomainStrategy {
|
||||
if x != nil {
|
||||
return x.DomainStrategy
|
||||
}
|
||||
return DeviceConfig_FORCE_IP
|
||||
}
|
||||
|
||||
func (x *DeviceConfig) GetIsClient() bool {
|
||||
if x != nil {
|
||||
return x.IsClient
|
||||
@@ -283,7 +220,7 @@ const file_proxy_wireguard_config_proto_rawDesc = "" +
|
||||
"\n" +
|
||||
"keep_alive\x18\x04 \x01(\tR\tkeepAlive\x12\x1f\n" +
|
||||
"\vallowed_ips\x18\x05 \x03(\tR\n" +
|
||||
"allowedIps\"\xee\x03\n" +
|
||||
"allowedIps\"\xb4\x02\n" +
|
||||
"\fDeviceConfig\x12\x1d\n" +
|
||||
"\n" +
|
||||
"secret_key\x18\x01 \x01(\tR\tsecretKey\x12\x1a\n" +
|
||||
@@ -291,20 +228,11 @@ const file_proxy_wireguard_config_proto_rawDesc = "" +
|
||||
"\x05peers\x18\x03 \x03(\v2 .xray.proxy.wireguard.PeerConfigR\x05peers\x120\n" +
|
||||
"\x05users\x18\x05 \x03(\v2\x1a.xray.common.protocol.UserR\x05users\x12\x10\n" +
|
||||
"\x03mtu\x18\x04 \x01(\x05R\x03mtu\x12\x1a\n" +
|
||||
"\breserved\x18\x06 \x01(\fR\breserved\x12Z\n" +
|
||||
"\x0fdomain_strategy\x18\a \x01(\x0e21.xray.proxy.wireguard.DeviceConfig.DomainStrategyR\x0edomainStrategy\x12\x1b\n" +
|
||||
"\breserved\x18\x06 \x01(\fR\breserved\x12\x1b\n" +
|
||||
"\tis_client\x18\b \x01(\bR\bisClient\x12\"\n" +
|
||||
"\rno_kernel_tun\x18\t \x01(\bR\vnoKernelTun\x12\x10\n" +
|
||||
"\x03DNS\x18\n" +
|
||||
" \x03(\tR\x03DNS\"\\\n" +
|
||||
"\x0eDomainStrategy\x12\f\n" +
|
||||
"\bFORCE_IP\x10\x00\x12\r\n" +
|
||||
"\tFORCE_IP4\x10\x01\x12\r\n" +
|
||||
"\tFORCE_IP6\x10\x02\x12\x0e\n" +
|
||||
"\n" +
|
||||
"FORCE_IP46\x10\x03\x12\x0e\n" +
|
||||
"\n" +
|
||||
"FORCE_IP64\x10\x04B^\n" +
|
||||
" \x03(\tR\x03DNSB^\n" +
|
||||
"\x18com.xray.proxy.wireguardP\x01Z)github.com/xtls/xray-core/proxy/wireguard\xaa\x02\x14Xray.Proxy.WireGuardb\x06proto3"
|
||||
|
||||
var (
|
||||
@@ -319,23 +247,20 @@ func file_proxy_wireguard_config_proto_rawDescGZIP() []byte {
|
||||
return file_proxy_wireguard_config_proto_rawDescData
|
||||
}
|
||||
|
||||
var file_proxy_wireguard_config_proto_enumTypes = make([]protoimpl.EnumInfo, 1)
|
||||
var file_proxy_wireguard_config_proto_msgTypes = make([]protoimpl.MessageInfo, 2)
|
||||
var file_proxy_wireguard_config_proto_goTypes = []any{
|
||||
(DeviceConfig_DomainStrategy)(0), // 0: xray.proxy.wireguard.DeviceConfig.DomainStrategy
|
||||
(*PeerConfig)(nil), // 1: xray.proxy.wireguard.PeerConfig
|
||||
(*DeviceConfig)(nil), // 2: xray.proxy.wireguard.DeviceConfig
|
||||
(*protocol.User)(nil), // 3: xray.common.protocol.User
|
||||
(*PeerConfig)(nil), // 0: xray.proxy.wireguard.PeerConfig
|
||||
(*DeviceConfig)(nil), // 1: xray.proxy.wireguard.DeviceConfig
|
||||
(*protocol.User)(nil), // 2: xray.common.protocol.User
|
||||
}
|
||||
var file_proxy_wireguard_config_proto_depIdxs = []int32{
|
||||
1, // 0: xray.proxy.wireguard.DeviceConfig.peers:type_name -> xray.proxy.wireguard.PeerConfig
|
||||
3, // 1: xray.proxy.wireguard.DeviceConfig.users:type_name -> xray.common.protocol.User
|
||||
0, // 2: xray.proxy.wireguard.DeviceConfig.domain_strategy:type_name -> xray.proxy.wireguard.DeviceConfig.DomainStrategy
|
||||
3, // [3:3] is the sub-list for method output_type
|
||||
3, // [3:3] is the sub-list for method input_type
|
||||
3, // [3:3] is the sub-list for extension type_name
|
||||
3, // [3:3] is the sub-list for extension extendee
|
||||
0, // [0:3] is the sub-list for field type_name
|
||||
0, // 0: xray.proxy.wireguard.DeviceConfig.peers:type_name -> xray.proxy.wireguard.PeerConfig
|
||||
2, // 1: xray.proxy.wireguard.DeviceConfig.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_wireguard_config_proto_init() }
|
||||
@@ -348,14 +273,13 @@ func file_proxy_wireguard_config_proto_init() {
|
||||
File: protoimpl.DescBuilder{
|
||||
GoPackagePath: reflect.TypeOf(x{}).PkgPath(),
|
||||
RawDescriptor: unsafe.Slice(unsafe.StringData(file_proxy_wireguard_config_proto_rawDesc), len(file_proxy_wireguard_config_proto_rawDesc)),
|
||||
NumEnums: 1,
|
||||
NumEnums: 0,
|
||||
NumMessages: 2,
|
||||
NumExtensions: 0,
|
||||
NumServices: 0,
|
||||
},
|
||||
GoTypes: file_proxy_wireguard_config_proto_goTypes,
|
||||
DependencyIndexes: file_proxy_wireguard_config_proto_depIdxs,
|
||||
EnumInfos: file_proxy_wireguard_config_proto_enumTypes,
|
||||
MessageInfos: file_proxy_wireguard_config_proto_msgTypes,
|
||||
}.Build()
|
||||
File_proxy_wireguard_config_proto = out.File
|
||||
|
||||
@@ -17,13 +17,6 @@ message PeerConfig {
|
||||
}
|
||||
|
||||
message DeviceConfig {
|
||||
enum DomainStrategy {
|
||||
FORCE_IP = 0;
|
||||
FORCE_IP4 = 1;
|
||||
FORCE_IP6 = 2;
|
||||
FORCE_IP46 = 3;
|
||||
FORCE_IP64 = 4;
|
||||
}
|
||||
string secret_key = 1;
|
||||
repeated string endpoint = 2;
|
||||
repeated PeerConfig peers = 3;
|
||||
@@ -31,7 +24,6 @@ message DeviceConfig {
|
||||
int32 mtu = 4;
|
||||
|
||||
bytes reserved = 6;
|
||||
DomainStrategy domain_strategy = 7;
|
||||
bool is_client = 8;
|
||||
bool no_kernel_tun = 9;
|
||||
repeated string DNS = 10;
|
||||
|
||||
+157
-14
@@ -15,6 +15,8 @@ import (
|
||||
"net"
|
||||
"net/netip"
|
||||
"os"
|
||||
"regexp"
|
||||
"strconv"
|
||||
"strings"
|
||||
"syscall"
|
||||
"time"
|
||||
@@ -42,6 +44,7 @@ type netTun struct {
|
||||
events chan tun.Event
|
||||
notifyHandle *channel.NotificationHandle
|
||||
incomingPacket chan *buffer.View
|
||||
closed chan struct{}
|
||||
mtu int
|
||||
dnsServers []netip.Addr
|
||||
hasV4, hasV6 bool
|
||||
@@ -58,6 +61,7 @@ func CreateNetTUN(localAddresses, dnsServers []netip.Addr, mtu int, handleLocal
|
||||
stack: stack.New(opts),
|
||||
events: make(chan tun.Event, 10),
|
||||
incomingPacket: make(chan *buffer.View),
|
||||
closed: make(chan struct{}),
|
||||
dnsServers: dnsServers,
|
||||
mtu: mtu,
|
||||
}
|
||||
@@ -124,8 +128,10 @@ func (tun *netTun) Events() <-chan tun.Event {
|
||||
}
|
||||
|
||||
func (tun *netTun) Read(buf [][]byte, sizes []int, offset int) (int, error) {
|
||||
view, ok := <-tun.incomingPacket
|
||||
if !ok {
|
||||
var view *buffer.View
|
||||
select {
|
||||
case view = <-tun.incomingPacket:
|
||||
case <-tun.closed:
|
||||
return 0, os.ErrClosed
|
||||
}
|
||||
|
||||
@@ -166,7 +172,10 @@ func (tun *netTun) WriteNotify() {
|
||||
view := pkt.ToView()
|
||||
pkt.DecRef()
|
||||
|
||||
tun.incomingPacket <- view
|
||||
select {
|
||||
case tun.incomingPacket <- view:
|
||||
case <-tun.closed:
|
||||
}
|
||||
}
|
||||
|
||||
func (tun *netTun) Close() error {
|
||||
@@ -179,8 +188,9 @@ func (tun *netTun) Close() error {
|
||||
close(tun.events)
|
||||
}
|
||||
|
||||
if tun.incomingPacket != nil {
|
||||
close(tun.incomingPacket)
|
||||
// we don't close incomingPacket, because WriteNotify may be mid-send on it (DNS lookup) and would panic.
|
||||
if tun.closed != nil {
|
||||
close(tun.closed)
|
||||
}
|
||||
|
||||
return nil
|
||||
@@ -219,6 +229,7 @@ type Net struct {
|
||||
DialUDPAddrPort func(laddr, raddr netip.AddrPort) (net.Conn, error)
|
||||
dnsServers []netip.Addr
|
||||
hasV4, hasV6 bool
|
||||
cache cache
|
||||
}
|
||||
|
||||
func convertToFullAddr(endpoint netip.AddrPort) (tcpip.FullAddress, tcpip.NetworkProtocolNumber) {
|
||||
@@ -246,9 +257,12 @@ var (
|
||||
errServerTemporarilyMisbehaving = errors.New("server misbehaving")
|
||||
errCanceled = errors.New("operation was canceled")
|
||||
errTimeout = errors.New("i/o timeout")
|
||||
errNumericPort = errors.New("port must be numeric")
|
||||
errNoSuitableAddress = errors.New("no suitable address found")
|
||||
errMissingAddress = errors.New("missing address")
|
||||
)
|
||||
|
||||
func (net *Net) LookupHost(host string) (addrs []net.IP, ttl uint32, err error) {
|
||||
func (net *Net) LookupHost(host string) (addrs []string, err error) {
|
||||
return net.LookupContextHost(context.Background(), host)
|
||||
}
|
||||
|
||||
@@ -567,9 +581,12 @@ func (tnet *Net) tryOneName(ctx context.Context, name string, qtype dnsmessage.T
|
||||
return dnsmessage.Parser{}, "", lastErr
|
||||
}
|
||||
|
||||
func (tnet *Net) LookupContextHost(ctx context.Context, host string) ([]net.IP, uint32, error) {
|
||||
func (tnet *Net) LookupContextHost(ctx context.Context, host string) ([]string, error) {
|
||||
if saddr := tnet.cache.LookupHost(host); saddr != nil {
|
||||
return saddr, nil
|
||||
}
|
||||
if host == "" || (!tnet.hasV6 && !tnet.hasV4) {
|
||||
return nil, 0, &net.DNSError{Err: errNoSuchHost.Error(), Name: host, IsNotFound: true}
|
||||
return nil, &net.DNSError{Err: errNoSuchHost.Error(), Name: host, IsNotFound: true}
|
||||
}
|
||||
zlen := len(host)
|
||||
if strings.IndexByte(host, ':') != -1 {
|
||||
@@ -578,11 +595,11 @@ func (tnet *Net) LookupContextHost(ctx context.Context, host string) ([]net.IP,
|
||||
}
|
||||
}
|
||||
if ip, err := netip.ParseAddr(host[:zlen]); err == nil {
|
||||
return []net.IP{ip.AsSlice()}, 0, nil
|
||||
return []string{ip.String()}, nil
|
||||
}
|
||||
|
||||
if !isDomainName(host) {
|
||||
return nil, 0, &net.DNSError{Err: errNoSuchHost.Error(), Name: host, IsNotFound: true}
|
||||
return nil, &net.DNSError{Err: errNoSuchHost.Error(), Name: host, IsNotFound: true}
|
||||
}
|
||||
type result struct {
|
||||
p dnsmessage.Parser
|
||||
@@ -683,11 +700,137 @@ func (tnet *Net) LookupContextHost(ctx context.Context, host string) ([]net.IP,
|
||||
}
|
||||
|
||||
if len(addrs) == 0 && lastErr != nil {
|
||||
return nil, 0, lastErr
|
||||
return nil, lastErr
|
||||
}
|
||||
ips := make([]net.IP, 0, len(addrs))
|
||||
saddrs := make([]string, 0, len(addrs))
|
||||
for _, ip := range addrs {
|
||||
ips = append(ips, ip.AsSlice())
|
||||
saddrs = append(saddrs, ip.String())
|
||||
}
|
||||
return ips, ttl, nil
|
||||
tnet.cache.Cache(host, saddrs, ttl)
|
||||
return saddrs, nil
|
||||
}
|
||||
|
||||
func partialDeadline(now, deadline time.Time, addrsRemaining int) (time.Time, error) {
|
||||
if deadline.IsZero() {
|
||||
return deadline, nil
|
||||
}
|
||||
timeRemaining := deadline.Sub(now)
|
||||
if timeRemaining <= 0 {
|
||||
return time.Time{}, errTimeout
|
||||
}
|
||||
timeout := timeRemaining / time.Duration(addrsRemaining)
|
||||
const saneMinimum = 2 * time.Second
|
||||
if timeout < saneMinimum {
|
||||
if timeRemaining < saneMinimum {
|
||||
timeout = timeRemaining
|
||||
} else {
|
||||
timeout = saneMinimum
|
||||
}
|
||||
}
|
||||
return now.Add(timeout), nil
|
||||
}
|
||||
|
||||
var protoSplitter = regexp.MustCompile(`^(tcp|udp|ping)(4|6)?$`)
|
||||
|
||||
func (tnet *Net) DialContext(ctx context.Context, network, address string) (net.Conn, error) {
|
||||
if ctx == nil {
|
||||
panic("nil context")
|
||||
}
|
||||
var acceptV4, acceptV6 bool
|
||||
matches := protoSplitter.FindStringSubmatch(network)
|
||||
if matches == nil {
|
||||
return nil, &net.OpError{Op: "dial", Err: net.UnknownNetworkError(network)}
|
||||
} else if len(matches[2]) == 0 {
|
||||
acceptV4 = true
|
||||
acceptV6 = true
|
||||
} else {
|
||||
acceptV4 = matches[2][0] == '4'
|
||||
acceptV6 = !acceptV4
|
||||
}
|
||||
var host string
|
||||
var port int
|
||||
if matches[1] == "ping" {
|
||||
host = address
|
||||
} else {
|
||||
var sport string
|
||||
var err error
|
||||
host, sport, err = net.SplitHostPort(address)
|
||||
if err != nil {
|
||||
return nil, &net.OpError{Op: "dial", Err: err}
|
||||
}
|
||||
port, err = strconv.Atoi(sport)
|
||||
if err != nil || port < 0 || port > 65535 {
|
||||
return nil, &net.OpError{Op: "dial", Err: errNumericPort}
|
||||
}
|
||||
}
|
||||
allAddr, err := tnet.LookupContextHost(ctx, host)
|
||||
if err != nil {
|
||||
return nil, &net.OpError{Op: "dial", Err: err}
|
||||
}
|
||||
var addrs []netip.AddrPort
|
||||
for _, addr := range allAddr {
|
||||
ip, err := netip.ParseAddr(addr)
|
||||
if err == nil && ((ip.Is4() && acceptV4) || (ip.Is6() && acceptV6)) {
|
||||
addrs = append(addrs, netip.AddrPortFrom(ip, uint16(port)))
|
||||
}
|
||||
}
|
||||
if len(addrs) == 0 && len(allAddr) != 0 {
|
||||
return nil, &net.OpError{Op: "dial", Err: errNoSuitableAddress}
|
||||
}
|
||||
|
||||
var firstErr error
|
||||
for i, addr := range addrs {
|
||||
select {
|
||||
case <-ctx.Done():
|
||||
err := ctx.Err()
|
||||
if err == context.Canceled {
|
||||
err = errCanceled
|
||||
} else if err == context.DeadlineExceeded {
|
||||
err = errTimeout
|
||||
}
|
||||
return nil, &net.OpError{Op: "dial", Err: err}
|
||||
default:
|
||||
}
|
||||
|
||||
dialCtx := ctx
|
||||
if deadline, hasDeadline := ctx.Deadline(); hasDeadline {
|
||||
partialDeadline, err := partialDeadline(time.Now(), deadline, len(addrs)-i)
|
||||
if err != nil {
|
||||
if firstErr == nil {
|
||||
firstErr = &net.OpError{Op: "dial", Err: err}
|
||||
}
|
||||
break
|
||||
}
|
||||
if partialDeadline.Before(deadline) {
|
||||
var cancel context.CancelFunc
|
||||
dialCtx, cancel = context.WithDeadline(ctx, partialDeadline)
|
||||
defer cancel()
|
||||
}
|
||||
}
|
||||
|
||||
var c net.Conn
|
||||
switch matches[1] {
|
||||
case "tcp":
|
||||
c, err = tnet.DialContextTCPAddrPort(dialCtx, addr)
|
||||
case "udp":
|
||||
c, err = tnet.DialUDPAddrPort(netip.AddrPort{}, addr)
|
||||
case "ping":
|
||||
err = errors.New("not support")
|
||||
// c, err = tnet.DialPingAddr(netip.Addr{}, addr.Addr())
|
||||
}
|
||||
if err == nil {
|
||||
return c, nil
|
||||
}
|
||||
if firstErr == nil {
|
||||
firstErr = err
|
||||
}
|
||||
}
|
||||
if firstErr == nil {
|
||||
firstErr = &net.OpError{Op: "dial", Err: errMissingAddress}
|
||||
}
|
||||
return nil, firstErr
|
||||
}
|
||||
|
||||
func (tnet *Net) Dial(network, address string) (net.Conn, error) {
|
||||
return tnet.DialContext(context.Background(), network, address)
|
||||
}
|
||||
|
||||
@@ -258,18 +258,16 @@ func (s *Server) Start() error {
|
||||
return errors.New("address is domain")
|
||||
}
|
||||
listenFunc := func() (net.PacketConn, error) {
|
||||
pktConn, err := internet.ListenSystemPacket(context.Background(), &net.UDPAddr{IP: s.src.Address.IP(), Port: int(s.src.Port)}, s.streamSettings.SocketSettings)
|
||||
var pktConn net.PacketConn
|
||||
var err error
|
||||
if s.streamSettings.FinalMask != nil {
|
||||
pktConn, err = s.streamSettings.FinalMask.ListenPacket(context.Background(), &net.UDPAddr{IP: s.src.Address.IP(), Port: int(s.src.Port)})
|
||||
} else {
|
||||
pktConn, err = internet.ListenSystemPacket(context.Background(), &net.UDPAddr{IP: s.src.Address.IP(), Port: int(s.src.Port)}, s.streamSettings.SocketSettings)
|
||||
}
|
||||
if err != nil {
|
||||
return nil, err
|
||||
}
|
||||
if s.streamSettings.UdpmaskManager != nil {
|
||||
newConn, err := s.streamSettings.UdpmaskManager.WrapPacketConnServer(pktConn)
|
||||
if err != nil {
|
||||
pktConn.Close()
|
||||
return nil, errors.New("mask err").Base(err)
|
||||
}
|
||||
pktConn = newConn
|
||||
}
|
||||
if s.uplinkCounter != nil || s.downlinkCounter != nil {
|
||||
pktConn = &PacketCounterConnection{
|
||||
PacketConn: pktConn,
|
||||
|
||||
@@ -0,0 +1,221 @@
|
||||
package scenarios
|
||||
|
||||
import (
|
||||
"crypto/rand"
|
||||
"encoding/base64"
|
||||
"testing"
|
||||
"time"
|
||||
|
||||
"github.com/sagernet/sing-shadowsocks/shadowaead_2022"
|
||||
"github.com/xtls/xray-core/app/log"
|
||||
"github.com/xtls/xray-core/app/proxyman"
|
||||
"github.com/xtls/xray-core/common"
|
||||
clog "github.com/xtls/xray-core/common/log"
|
||||
"github.com/xtls/xray-core/common/net"
|
||||
"github.com/xtls/xray-core/common/serial"
|
||||
"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/shadowsocks_2022"
|
||||
"github.com/xtls/xray-core/testing/servers/tcp"
|
||||
"github.com/xtls/xray-core/testing/servers/udp"
|
||||
"golang.org/x/sync/errgroup"
|
||||
)
|
||||
|
||||
func TestShadowsocks2022Tcp(t *testing.T) {
|
||||
for _, method := range shadowaead_2022.List {
|
||||
password := make([]byte, 32)
|
||||
rand.Read(password)
|
||||
t.Run(method, func(t *testing.T) {
|
||||
testShadowsocks2022Tcp(t, method, base64.StdEncoding.EncodeToString(password))
|
||||
})
|
||||
}
|
||||
}
|
||||
|
||||
func TestShadowsocks2022UdpAES128(t *testing.T) {
|
||||
password := make([]byte, 32)
|
||||
rand.Read(password)
|
||||
testShadowsocks2022Udp(t, shadowaead_2022.List[0], base64.StdEncoding.EncodeToString(password))
|
||||
}
|
||||
|
||||
func TestShadowsocks2022UdpAES256(t *testing.T) {
|
||||
password := make([]byte, 32)
|
||||
rand.Read(password)
|
||||
testShadowsocks2022Udp(t, shadowaead_2022.List[1], base64.StdEncoding.EncodeToString(password))
|
||||
}
|
||||
|
||||
func TestShadowsocks2022UdpChacha(t *testing.T) {
|
||||
password := make([]byte, 32)
|
||||
rand.Read(password)
|
||||
testShadowsocks2022Udp(t, shadowaead_2022.List[2], base64.StdEncoding.EncodeToString(password))
|
||||
}
|
||||
|
||||
func testShadowsocks2022Tcp(t *testing.T, method string, password string) {
|
||||
tcpServer := tcp.Server{
|
||||
MsgProcessor: xor,
|
||||
}
|
||||
dest, err := tcpServer.Start()
|
||||
common.Must(err)
|
||||
defer tcpServer.Close()
|
||||
|
||||
serverPort := tcp.PickPort()
|
||||
serverConfig := &core.Config{
|
||||
App: []*serial.TypedMessage{
|
||||
serial.ToTypedMessage(&log.Config{
|
||||
ErrorLogLevel: clog.Severity_Debug,
|
||||
ErrorLogType: log.LogType_Console,
|
||||
}),
|
||||
},
|
||||
Inbound: []*core.InboundHandlerConfig{
|
||||
{
|
||||
ReceiverSettings: serial.ToTypedMessage(&proxyman.ReceiverConfig{
|
||||
PortList: &net.PortList{Range: []*net.PortRange{net.SinglePortRange(serverPort)}},
|
||||
Listen: net.NewIPOrDomain(net.LocalHostIP),
|
||||
}),
|
||||
ProxySettings: serial.ToTypedMessage(&shadowsocks_2022.ServerConfig{
|
||||
Method: method,
|
||||
Key: password,
|
||||
Network: []net.Network{net.Network_TCP},
|
||||
}),
|
||||
},
|
||||
},
|
||||
Outbound: []*core.OutboundHandlerConfig{
|
||||
{
|
||||
ProxySettings: serial.ToTypedMessage(&freedom.Config{
|
||||
FinalRules: []*freedom.FinalRuleConfig{{Action: freedom.RuleAction_Allow}},
|
||||
}),
|
||||
},
|
||||
},
|
||||
}
|
||||
|
||||
clientPort := tcp.PickPort()
|
||||
clientConfig := &core.Config{
|
||||
App: []*serial.TypedMessage{
|
||||
serial.ToTypedMessage(&log.Config{
|
||||
ErrorLogLevel: clog.Severity_Debug,
|
||||
ErrorLogType: log.LogType_Console,
|
||||
}),
|
||||
},
|
||||
Inbound: []*core.InboundHandlerConfig{
|
||||
{
|
||||
ReceiverSettings: serial.ToTypedMessage(&proxyman.ReceiverConfig{
|
||||
PortList: &net.PortList{Range: []*net.PortRange{net.SinglePortRange(clientPort)}},
|
||||
Listen: net.NewIPOrDomain(net.LocalHostIP),
|
||||
}),
|
||||
ProxySettings: serial.ToTypedMessage(&dokodemo.Config{
|
||||
RewriteAddress: net.NewIPOrDomain(dest.Address),
|
||||
RewritePort: uint32(dest.Port),
|
||||
AllowedNetworks: []net.Network{net.Network_TCP},
|
||||
}),
|
||||
},
|
||||
},
|
||||
Outbound: []*core.OutboundHandlerConfig{
|
||||
{
|
||||
ProxySettings: serial.ToTypedMessage(&shadowsocks_2022.ClientConfig{
|
||||
Address: net.NewIPOrDomain(net.LocalHostIP),
|
||||
Port: uint32(serverPort),
|
||||
Method: method,
|
||||
Key: password,
|
||||
}),
|
||||
},
|
||||
},
|
||||
}
|
||||
|
||||
servers, err := InitializeServerConfigs(serverConfig, clientConfig)
|
||||
common.Must(err)
|
||||
defer CloseAllServers(servers)
|
||||
|
||||
var errGroup errgroup.Group
|
||||
for range 3 {
|
||||
errGroup.Go(testTCPConn(clientPort, 10240*1024, time.Second*20))
|
||||
}
|
||||
|
||||
if err := errGroup.Wait(); err != nil {
|
||||
t.Error(err)
|
||||
}
|
||||
}
|
||||
|
||||
func testShadowsocks2022Udp(t *testing.T, method string, password string) {
|
||||
udpServer := udp.Server{
|
||||
MsgProcessor: xor,
|
||||
}
|
||||
udpDest, err := udpServer.Start()
|
||||
common.Must(err)
|
||||
defer udpServer.Close()
|
||||
|
||||
serverPort := udp.PickPort()
|
||||
serverConfig := &core.Config{
|
||||
App: []*serial.TypedMessage{
|
||||
serial.ToTypedMessage(&log.Config{
|
||||
ErrorLogLevel: clog.Severity_Debug,
|
||||
ErrorLogType: log.LogType_Console,
|
||||
}),
|
||||
},
|
||||
Inbound: []*core.InboundHandlerConfig{
|
||||
{
|
||||
ReceiverSettings: serial.ToTypedMessage(&proxyman.ReceiverConfig{
|
||||
PortList: &net.PortList{Range: []*net.PortRange{net.SinglePortRange(serverPort)}},
|
||||
Listen: net.NewIPOrDomain(net.LocalHostIP),
|
||||
}),
|
||||
ProxySettings: serial.ToTypedMessage(&shadowsocks_2022.ServerConfig{
|
||||
Method: method,
|
||||
Key: password,
|
||||
Network: []net.Network{net.Network_UDP},
|
||||
}),
|
||||
},
|
||||
},
|
||||
Outbound: []*core.OutboundHandlerConfig{
|
||||
{
|
||||
ProxySettings: serial.ToTypedMessage(&freedom.Config{
|
||||
FinalRules: []*freedom.FinalRuleConfig{{Action: freedom.RuleAction_Allow}},
|
||||
}),
|
||||
},
|
||||
},
|
||||
}
|
||||
|
||||
udpClientPort := udp.PickPort()
|
||||
clientConfig := &core.Config{
|
||||
App: []*serial.TypedMessage{
|
||||
serial.ToTypedMessage(&log.Config{
|
||||
ErrorLogLevel: clog.Severity_Debug,
|
||||
ErrorLogType: log.LogType_Console,
|
||||
}),
|
||||
},
|
||||
Inbound: []*core.InboundHandlerConfig{
|
||||
{
|
||||
ReceiverSettings: serial.ToTypedMessage(&proxyman.ReceiverConfig{
|
||||
PortList: &net.PortList{Range: []*net.PortRange{net.SinglePortRange(udpClientPort)}},
|
||||
Listen: net.NewIPOrDomain(net.LocalHostIP),
|
||||
}),
|
||||
ProxySettings: serial.ToTypedMessage(&dokodemo.Config{
|
||||
RewriteAddress: net.NewIPOrDomain(udpDest.Address),
|
||||
RewritePort: uint32(udpDest.Port),
|
||||
AllowedNetworks: []net.Network{net.Network_UDP},
|
||||
}),
|
||||
},
|
||||
},
|
||||
Outbound: []*core.OutboundHandlerConfig{
|
||||
{
|
||||
ProxySettings: serial.ToTypedMessage(&shadowsocks_2022.ClientConfig{
|
||||
Address: net.NewIPOrDomain(net.LocalHostIP),
|
||||
Port: uint32(serverPort),
|
||||
Method: method,
|
||||
Key: password,
|
||||
}),
|
||||
},
|
||||
},
|
||||
}
|
||||
|
||||
servers, err := InitializeServerConfigs(serverConfig, clientConfig)
|
||||
common.Must(err)
|
||||
defer CloseAllServers(servers)
|
||||
|
||||
var errGroup errgroup.Group
|
||||
for range 3 {
|
||||
errGroup.Go(testUDPConn(udpClientPort, 1024, time.Second*5))
|
||||
}
|
||||
|
||||
if err := errGroup.Wait(); err != nil {
|
||||
t.Error(err)
|
||||
}
|
||||
}
|
||||
@@ -65,6 +65,7 @@ func TestWireguard(t *testing.T) {
|
||||
ProxySettings: serial.ToTypedMessage(&freedom.Config{
|
||||
FinalRules: []*freedom.FinalRuleConfig{{Action: freedom.RuleAction_Allow}},
|
||||
}),
|
||||
SenderSettings: serial.ToTypedMessage(&proxyman.SenderConfig{}),
|
||||
},
|
||||
},
|
||||
}
|
||||
@@ -104,6 +105,7 @@ func TestWireguard(t *testing.T) {
|
||||
AllowedIps: []string{"0.0.0.0/0", "::0/0"},
|
||||
}},
|
||||
}),
|
||||
SenderSettings: serial.ToTypedMessage(&proxyman.SenderConfig{}),
|
||||
},
|
||||
},
|
||||
}
|
||||
|
||||
@@ -2,103 +2,291 @@ package finalmask
|
||||
|
||||
import (
|
||||
"context"
|
||||
"net"
|
||||
"fmt"
|
||||
"slices"
|
||||
|
||||
"github.com/xtls/xray-core/common/buf"
|
||||
"github.com/xtls/xray-core/common/errors"
|
||||
"github.com/xtls/xray-core/common/net"
|
||||
)
|
||||
|
||||
type Udpmask interface {
|
||||
WrapPacketConnClient(raw net.PacketConn, level int, levelCount int) (net.PacketConn, error)
|
||||
WrapPacketConnServer(raw net.PacketConn, level int, levelCount int) (net.PacketConn, error)
|
||||
type Dialer struct {
|
||||
DialTCP func(net.Destination) (net.Conn, error)
|
||||
DialUDP func(net.Destination) (net.Conn, error)
|
||||
}
|
||||
|
||||
type UdpmaskManager struct {
|
||||
udpmasks []Udpmask
|
||||
type ListenConfig struct {
|
||||
Listen func(net.Addr) (net.Listener, error)
|
||||
ListenPacket func(net.Addr) (net.PacketConn, error)
|
||||
}
|
||||
|
||||
func NewUdpmaskManager(udpmasks []Udpmask) *UdpmaskManager {
|
||||
slices.Reverse(udpmasks)
|
||||
return &UdpmaskManager{udpmasks: udpmasks}
|
||||
type TCPMask interface {
|
||||
WrapConnClient(net.Conn, *net.Destination, *Dialer) (net.Conn, error)
|
||||
WrapConnServer(net.Conn) (net.Conn, error)
|
||||
// Listen(net.Listener) (net.Listener, error)
|
||||
}
|
||||
|
||||
func (m *UdpmaskManager) WrapPacketConnClient(raw net.PacketConn) (net.PacketConn, error) {
|
||||
var sizes []int
|
||||
var conns []net.PacketConn
|
||||
for i, mask := range m.udpmasks {
|
||||
if _, ok := mask.(headerConn); ok {
|
||||
conn, err := mask.WrapPacketConnClient(nil, i, len(m.udpmasks)-1)
|
||||
if err != nil {
|
||||
return nil, err
|
||||
}
|
||||
sizes = append(sizes, conn.(headerSize).Size())
|
||||
conns = append(conns, conn)
|
||||
} else {
|
||||
if len(conns) > 0 {
|
||||
raw = &headerManagerConn{sizes: sizes, conns: conns, PacketConn: raw}
|
||||
sizes = nil
|
||||
conns = nil
|
||||
}
|
||||
var err error
|
||||
raw, err = mask.WrapPacketConnClient(raw, i, len(m.udpmasks)-1)
|
||||
if err != nil {
|
||||
return nil, err
|
||||
type UDPMask interface {
|
||||
WrapPacketConnClient(net.PacketConn, *net.Destination, *Dialer) (net.PacketConn, error)
|
||||
WrapPacketConnServer(net.PacketConn, net.Addr, *ListenConfig) (net.PacketConn, error)
|
||||
}
|
||||
|
||||
type FinalMask struct {
|
||||
tcpMasks []TCPMask
|
||||
udpMasks []UDPMask
|
||||
dialTCP func(context.Context, net.Destination) (net.Conn, error)
|
||||
listen func(context.Context, net.Addr) (net.Listener, error)
|
||||
dialUDP func(context.Context, net.Destination) (net.PacketConn, net.Addr, error)
|
||||
listenPacket func(context.Context, net.Addr) (net.PacketConn, error)
|
||||
}
|
||||
|
||||
func NewFinalMask(tcpMasks []TCPMask, udpMasks []UDPMask, dialTCP func(context.Context, net.Destination) (net.Conn, error), listen func(context.Context, net.Addr) (net.Listener, error), dialUDP func(context.Context, net.Destination) (net.PacketConn, net.Addr, error), listenPacket func(context.Context, net.Addr) (net.PacketConn, error)) *FinalMask {
|
||||
slices.Reverse(tcpMasks)
|
||||
slices.Reverse(udpMasks)
|
||||
return &FinalMask{
|
||||
tcpMasks: tcpMasks,
|
||||
udpMasks: udpMasks,
|
||||
dialTCP: dialTCP,
|
||||
dialUDP: dialUDP,
|
||||
listen: listen,
|
||||
listenPacket: listenPacket,
|
||||
}
|
||||
}
|
||||
|
||||
func (fm *FinalMask) DialTCP(ctx context.Context, dest net.Destination) (net.Conn, error) {
|
||||
if len(fm.tcpMasks) == 0 {
|
||||
return fm.dialTCP(ctx, dest)
|
||||
}
|
||||
for i := range fm.tcpMasks {
|
||||
if i > 0 {
|
||||
if _, ok := fm.tcpMasks[i].(interface{ HandleDial() }); ok {
|
||||
return nil, fmt.Errorf("incorrect index: %d %T", i, fm.tcpMasks[i])
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
if len(conns) > 0 {
|
||||
raw = &headerManagerConn{sizes: sizes, conns: conns, PacketConn: raw}
|
||||
sizes = nil
|
||||
conns = nil
|
||||
var conn net.Conn
|
||||
var err error
|
||||
if _, ok := fm.tcpMasks[0].(interface{ HandleDial() }); !ok {
|
||||
conn, err = fm.dialTCP(ctx, dest)
|
||||
if err != nil {
|
||||
return nil, err
|
||||
}
|
||||
}
|
||||
return raw, nil
|
||||
dialer := &Dialer{
|
||||
DialTCP: func(dest net.Destination) (net.Conn, error) {
|
||||
return fm.dialTCP(ctx, dest)
|
||||
},
|
||||
DialUDP: func(dest net.Destination) (net.Conn, error) {
|
||||
conn, addr, err := fm.dialUDP(ctx, dest)
|
||||
if err != nil {
|
||||
return nil, err
|
||||
}
|
||||
return &PacketConnWrapper{PacketConn: conn, udpAddr: addr}, err
|
||||
},
|
||||
}
|
||||
for i := range fm.tcpMasks {
|
||||
var newConn net.Conn
|
||||
newConn, err = fm.tcpMasks[i].WrapConnClient(conn, &dest, dialer)
|
||||
if err != nil {
|
||||
_ = conn.Close()
|
||||
return nil, err
|
||||
}
|
||||
conn = newConn
|
||||
}
|
||||
return conn, nil
|
||||
}
|
||||
|
||||
func (m *UdpmaskManager) WrapPacketConnServer(raw net.PacketConn) (net.PacketConn, error) {
|
||||
var sizes []int
|
||||
var conns []net.PacketConn
|
||||
for i, mask := range m.udpmasks {
|
||||
if _, ok := mask.(headerConn); ok {
|
||||
conn, err := mask.WrapPacketConnServer(nil, i, len(m.udpmasks)-1)
|
||||
if err != nil {
|
||||
return nil, err
|
||||
func (fm *FinalMask) Listen(ctx context.Context, addr net.Addr) (net.Listener, error) {
|
||||
if len(fm.tcpMasks) == 0 {
|
||||
return fm.listen(ctx, addr)
|
||||
}
|
||||
off := 0
|
||||
listener, err := fm.listen(ctx, addr)
|
||||
if err != nil {
|
||||
return nil, err
|
||||
}
|
||||
for i := range fm.tcpMasks {
|
||||
if _, ok := fm.tcpMasks[i].(interface {
|
||||
Listen(net.Listener) (net.Listener, error)
|
||||
}); ok {
|
||||
if i-off == 0 {
|
||||
l, err := fm.tcpMasks[i].(interface {
|
||||
Listen(net.Listener) (net.Listener, error)
|
||||
}).Listen(listener)
|
||||
if err != nil {
|
||||
listener.Close()
|
||||
return nil, err
|
||||
}
|
||||
listener = l
|
||||
} else {
|
||||
l, err := fm.tcpMasks[i].(interface {
|
||||
Listen(net.Listener) (net.Listener, error)
|
||||
}).Listen(&TCPListener{Listener: listener, tcpMasks: fm.tcpMasks[off:i]})
|
||||
if err != nil {
|
||||
listener.Close()
|
||||
return nil, err
|
||||
}
|
||||
listener = l
|
||||
}
|
||||
sizes = append(sizes, conn.(headerSize).Size())
|
||||
conns = append(conns, conn)
|
||||
} else {
|
||||
if len(conns) > 0 {
|
||||
raw = &headerManagerConn{sizes: sizes, conns: conns, PacketConn: raw}
|
||||
sizes = nil
|
||||
conns = nil
|
||||
}
|
||||
var err error
|
||||
raw, err = mask.WrapPacketConnServer(raw, i, len(m.udpmasks)-1)
|
||||
if err != nil {
|
||||
return nil, err
|
||||
off = i + 1
|
||||
}
|
||||
}
|
||||
if off < len(fm.tcpMasks) {
|
||||
return &TCPListener{Listener: listener, tcpMasks: fm.tcpMasks[off:]}, nil
|
||||
}
|
||||
return listener, nil
|
||||
}
|
||||
|
||||
func (fm *FinalMask) DialUDP(ctx context.Context, dest net.Destination) (net.Conn, error) {
|
||||
if len(fm.udpMasks) == 0 {
|
||||
conn, addr, err := fm.dialUDP(ctx, dest)
|
||||
if err != nil {
|
||||
return nil, err
|
||||
}
|
||||
return &PacketConnWrapper{PacketConn: conn, udpAddr: addr}, nil
|
||||
}
|
||||
for i := range fm.udpMasks {
|
||||
if i > 0 {
|
||||
if _, ok := fm.udpMasks[i].(interface{ HandleDial() }); ok {
|
||||
return nil, fmt.Errorf("incorrect index: %d %T", i, fm.udpMasks[i])
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
var conn net.PacketConn
|
||||
var addr net.Addr
|
||||
var err error
|
||||
if _, ok := fm.udpMasks[0].(interface{ HandleDial() }); !ok {
|
||||
conn, addr, err = fm.dialUDP(ctx, dest)
|
||||
if err != nil {
|
||||
return nil, err
|
||||
}
|
||||
}
|
||||
dialer := &Dialer{
|
||||
DialTCP: func(dest net.Destination) (net.Conn, error) {
|
||||
return fm.dialTCP(ctx, dest)
|
||||
},
|
||||
DialUDP: func(dest net.Destination) (net.Conn, error) {
|
||||
conn, addr, err := fm.dialUDP(ctx, dest)
|
||||
if err != nil {
|
||||
return nil, err
|
||||
}
|
||||
return &PacketConnWrapper{PacketConn: conn, udpAddr: addr}, err
|
||||
},
|
||||
}
|
||||
var sizes []int
|
||||
var conns []net.PacketConn
|
||||
for i := range fm.udpMasks {
|
||||
var newConn net.PacketConn
|
||||
if _, ok := fm.udpMasks[i].(interface{ HeaderConn() }); ok {
|
||||
newConn, err = fm.udpMasks[i].WrapPacketConnClient(nil, nil, nil)
|
||||
if err != nil {
|
||||
_ = conn.Close()
|
||||
return nil, err
|
||||
}
|
||||
sizes = append(sizes, newConn.(interface{ Size() int }).Size())
|
||||
conns = append(conns, newConn)
|
||||
} else {
|
||||
if len(conns) > 0 {
|
||||
conn = &headerManagerConn{PacketConn: conn, sizes: sizes, conns: conns}
|
||||
sizes = nil
|
||||
conns = nil
|
||||
}
|
||||
newConn, err = fm.udpMasks[i].WrapPacketConnClient(conn, &dest, dialer)
|
||||
if err != nil {
|
||||
_ = conn.Close()
|
||||
return nil, err
|
||||
}
|
||||
conn = newConn
|
||||
}
|
||||
}
|
||||
if len(conns) > 0 {
|
||||
raw = &headerManagerConn{sizes: sizes, conns: conns, PacketConn: raw}
|
||||
conn = &headerManagerConn{PacketConn: conn, sizes: sizes, conns: conns}
|
||||
sizes = nil
|
||||
conns = nil
|
||||
}
|
||||
return raw, nil
|
||||
if addr == nil {
|
||||
addr = &net.UDPAddr{IP: []byte{0, 0, 0, 0}}
|
||||
}
|
||||
return &PacketConnWrapper{PacketConn: conn, udpAddr: addr}, nil
|
||||
}
|
||||
|
||||
func (fm *FinalMask) ListenPacket(ctx context.Context, addr net.Addr) (net.PacketConn, error) {
|
||||
if len(fm.udpMasks) == 0 {
|
||||
return fm.listenPacket(ctx, addr)
|
||||
}
|
||||
for i := range fm.udpMasks {
|
||||
if i > 0 {
|
||||
if _, ok := fm.udpMasks[i].(interface{ HandleListen() }); ok {
|
||||
return nil, fmt.Errorf("incorrect index: %d %T", i, fm.udpMasks[i])
|
||||
}
|
||||
}
|
||||
}
|
||||
var conn net.PacketConn
|
||||
var err error
|
||||
if _, ok := fm.udpMasks[0].(interface{ HandleListen() }); !ok {
|
||||
conn, err = fm.listenPacket(ctx, addr)
|
||||
if err != nil {
|
||||
return nil, err
|
||||
}
|
||||
}
|
||||
lc := &ListenConfig{
|
||||
Listen: func(addr net.Addr) (net.Listener, error) { return fm.listen(ctx, addr) },
|
||||
ListenPacket: func(addr net.Addr) (net.PacketConn, error) { return fm.listenPacket(ctx, addr) },
|
||||
}
|
||||
var sizes []int
|
||||
var conns []net.PacketConn
|
||||
for i := range fm.udpMasks {
|
||||
var newConn net.PacketConn
|
||||
if _, ok := fm.udpMasks[i].(interface{ HeaderConn() }); ok {
|
||||
newConn, err = fm.udpMasks[i].WrapPacketConnServer(nil, nil, nil)
|
||||
if err != nil {
|
||||
_ = conn.Close()
|
||||
return nil, err
|
||||
}
|
||||
sizes = append(sizes, newConn.(interface{ Size() int }).Size())
|
||||
conns = append(conns, newConn)
|
||||
} else {
|
||||
if len(conns) > 0 {
|
||||
conn = &headerManagerConn{PacketConn: conn, sizes: sizes, conns: conns}
|
||||
sizes = nil
|
||||
conns = nil
|
||||
}
|
||||
newConn, err = fm.udpMasks[i].WrapPacketConnServer(conn, addr, lc)
|
||||
if err != nil {
|
||||
_ = conn.Close()
|
||||
return nil, err
|
||||
}
|
||||
conn = newConn
|
||||
}
|
||||
}
|
||||
if len(conns) > 0 {
|
||||
conn = &headerManagerConn{PacketConn: conn, sizes: sizes, conns: conns}
|
||||
sizes = nil
|
||||
conns = nil
|
||||
}
|
||||
return conn, nil
|
||||
}
|
||||
|
||||
const (
|
||||
UDPSize = 4096
|
||||
)
|
||||
|
||||
type headerConn interface {
|
||||
HeaderConn()
|
||||
type PacketConnWrapper struct {
|
||||
net.PacketConn
|
||||
udpAddr net.Addr
|
||||
}
|
||||
|
||||
type headerSize interface {
|
||||
Size() int
|
||||
func (c *PacketConnWrapper) RemoteAddr() net.Addr {
|
||||
return c.udpAddr
|
||||
}
|
||||
|
||||
func (c *PacketConnWrapper) Read(b []byte) (n int, err error) {
|
||||
n, _, err = c.PacketConn.ReadFrom(b)
|
||||
return
|
||||
}
|
||||
|
||||
func (c *PacketConnWrapper) Write(b []byte) (n int, err error) {
|
||||
return c.PacketConn.WriteTo(b, c.udpAddr)
|
||||
}
|
||||
|
||||
type headerManagerConn struct {
|
||||
@@ -191,72 +379,27 @@ func (c *headerManagerConn) WriteTo(p []byte, addr net.Addr) (n int, err error)
|
||||
return len(p), nil
|
||||
}
|
||||
|
||||
type Tcpmask interface {
|
||||
WrapConnClient(net.Conn) (net.Conn, error)
|
||||
WrapConnServer(net.Conn) (net.Conn, error)
|
||||
}
|
||||
|
||||
type TcpmaskManager struct {
|
||||
tcpmasks []Tcpmask
|
||||
}
|
||||
|
||||
func NewTcpmaskManager(tcpmasks []Tcpmask) *TcpmaskManager {
|
||||
slices.Reverse(tcpmasks)
|
||||
return &TcpmaskManager{tcpmasks: tcpmasks}
|
||||
}
|
||||
|
||||
func (m *TcpmaskManager) WrapConnClient(raw net.Conn) (net.Conn, error) {
|
||||
var err error
|
||||
for _, mask := range m.tcpmasks {
|
||||
raw, err = mask.WrapConnClient(raw)
|
||||
if err != nil {
|
||||
return nil, err
|
||||
}
|
||||
}
|
||||
return raw, nil
|
||||
}
|
||||
|
||||
func (m *TcpmaskManager) WrapConnServer(raw net.Conn) (net.Conn, error) {
|
||||
var err error
|
||||
for _, mask := range m.tcpmasks {
|
||||
raw, err = mask.WrapConnServer(raw)
|
||||
if err != nil {
|
||||
return nil, err
|
||||
}
|
||||
}
|
||||
return raw, nil
|
||||
}
|
||||
|
||||
func (m *TcpmaskManager) WrapListener(l net.Listener) (net.Listener, error) {
|
||||
return NewTcpListener(m, l)
|
||||
}
|
||||
|
||||
type tcpListener struct {
|
||||
m *TcpmaskManager
|
||||
type TCPListener struct {
|
||||
net.Listener
|
||||
tcpMasks []TCPMask
|
||||
}
|
||||
|
||||
func NewTcpListener(m *TcpmaskManager, l net.Listener) (net.Listener, error) {
|
||||
return &tcpListener{
|
||||
m: m,
|
||||
Listener: l,
|
||||
}, nil
|
||||
}
|
||||
|
||||
func (l *tcpListener) Accept() (net.Conn, error) {
|
||||
func (l *TCPListener) Accept() (net.Conn, error) {
|
||||
conn, err := l.Listener.Accept()
|
||||
if err != nil {
|
||||
return conn, err
|
||||
}
|
||||
|
||||
newConn, err := l.m.WrapConnServer(conn)
|
||||
if err != nil {
|
||||
errors.LogDebugInner(context.Background(), err, "mask err")
|
||||
_ = conn.Close()
|
||||
return nil, err
|
||||
for i := range l.tcpMasks {
|
||||
var newConn net.Conn
|
||||
newConn, err = l.tcpMasks[i].WrapConnServer(conn)
|
||||
if err != nil {
|
||||
_ = conn.Close()
|
||||
return nil, err
|
||||
}
|
||||
conn = newConn
|
||||
}
|
||||
|
||||
return newConn, nil
|
||||
return conn, nil
|
||||
}
|
||||
|
||||
type TcpMaskConn interface {
|
||||
|
||||
@@ -1,11 +1,14 @@
|
||||
package fragment
|
||||
|
||||
import "net"
|
||||
import (
|
||||
"github.com/xtls/xray-core/common/net"
|
||||
"github.com/xtls/xray-core/transport/internet/finalmask"
|
||||
)
|
||||
|
||||
func (c *Config) WrapConnClient(raw net.Conn) (net.Conn, error) {
|
||||
return NewConnClient(c, raw, false)
|
||||
func (c *Config) WrapConnClient(conn net.Conn, dest *net.Destination, dialer *finalmask.Dialer) (net.Conn, error) {
|
||||
return NewConnClient(c, conn, false)
|
||||
}
|
||||
|
||||
func (c *Config) WrapConnServer(raw net.Conn) (net.Conn, error) {
|
||||
return NewConnServer(c, raw, true)
|
||||
func (c *Config) WrapConnServer(conn net.Conn) (net.Conn, error) {
|
||||
return NewConnServer(c, conn, true)
|
||||
}
|
||||
|
||||
@@ -1,29 +1,30 @@
|
||||
package custom
|
||||
|
||||
import (
|
||||
"net"
|
||||
"github.com/xtls/xray-core/common/net"
|
||||
"github.com/xtls/xray-core/transport/internet/finalmask"
|
||||
)
|
||||
|
||||
func (c *TCPConfig) WrapConnClient(raw net.Conn) (net.Conn, error) {
|
||||
return NewConnClientTCP(c, raw)
|
||||
func (c *TCPConfig) WrapConnClient(conn net.Conn, dest *net.Destination, dialer *finalmask.Dialer) (net.Conn, error) {
|
||||
return NewConnClientTCP(c, conn)
|
||||
}
|
||||
|
||||
func (c *TCPConfig) WrapConnServer(raw net.Conn) (net.Conn, error) {
|
||||
return NewConnServerTCP(c, raw)
|
||||
func (c *TCPConfig) WrapConnServer(conn net.Conn) (net.Conn, error) {
|
||||
return NewConnServerTCP(c, conn)
|
||||
}
|
||||
|
||||
func (c *UDPConfig) WrapPacketConnClient(raw net.PacketConn, level int, levelCount int) (net.PacketConn, error) {
|
||||
return NewConnClientUDP(c, raw)
|
||||
func (c *UDPConfig) WrapPacketConnClient(conn net.PacketConn, dest *net.Destination, dialer *finalmask.Dialer) (net.PacketConn, error) {
|
||||
return NewConnClientUDP(c, conn)
|
||||
}
|
||||
|
||||
func (c *UDPConfig) WrapPacketConnServer(raw net.PacketConn, level int, levelCount int) (net.PacketConn, error) {
|
||||
return NewConnServerUDP(c, raw)
|
||||
func (c *UDPConfig) WrapPacketConnServer(conn net.PacketConn, addr net.Addr, lc *finalmask.ListenConfig) (net.PacketConn, error) {
|
||||
return NewConnServerUDP(c, conn)
|
||||
}
|
||||
|
||||
func (c *UDPStandaloneConfig) WrapPacketConnClient(raw net.PacketConn, level int, levelCount int) (net.PacketConn, error) {
|
||||
return NewConnClientUDPStandalone(c, raw)
|
||||
func (c *UDPStandaloneConfig) WrapPacketConnClient(conn net.PacketConn, dest *net.Destination, dialer *finalmask.Dialer) (net.PacketConn, error) {
|
||||
return NewConnClientUDPStandalone(c, conn)
|
||||
}
|
||||
|
||||
func (c *UDPStandaloneConfig) WrapPacketConnServer(raw net.PacketConn, level int, levelCount int) (net.PacketConn, error) {
|
||||
return NewConnServerUDPStandalone(c, raw)
|
||||
func (c *UDPStandaloneConfig) WrapPacketConnServer(conn net.PacketConn, addr net.Addr, lc *finalmask.ListenConfig) (net.PacketConn, error) {
|
||||
return NewConnServerUDPStandalone(c, conn)
|
||||
}
|
||||
|
||||
@@ -9,8 +9,6 @@ import (
|
||||
"strings"
|
||||
"testing"
|
||||
"time"
|
||||
|
||||
"github.com/xtls/xray-core/transport/internet/finalmask"
|
||||
)
|
||||
|
||||
func TestMetadataEvaluatorRejectsUnknownName(t *testing.T) {
|
||||
@@ -156,7 +154,7 @@ func TestMetadataUDPStandaloneWriteUsesRemotePort(t *testing.T) {
|
||||
}
|
||||
defer serverRaw.Close()
|
||||
|
||||
client, err := finalmask.NewUdpmaskManager([]finalmask.Udpmask{cfg}).WrapPacketConnClient(clientRaw)
|
||||
client, err := cfg.WrapPacketConnClient(clientRaw, nil, nil)
|
||||
if err != nil {
|
||||
t.Fatal(err)
|
||||
}
|
||||
@@ -301,7 +299,7 @@ func TestMetadataTCPHandshakeUsesEndpointPorts(t *testing.T) {
|
||||
}
|
||||
defer serverRaw.Close()
|
||||
|
||||
client, err := clientCfg.WrapConnClient(clientRaw)
|
||||
client, err := clientCfg.WrapConnClient(clientRaw, nil, nil)
|
||||
if err != nil {
|
||||
t.Fatal(err)
|
||||
}
|
||||
|
||||
@@ -5,8 +5,6 @@ import (
|
||||
"net"
|
||||
"testing"
|
||||
"time"
|
||||
|
||||
"github.com/xtls/xray-core/transport/internet/finalmask"
|
||||
)
|
||||
|
||||
func mustSendRecvUDP(t *testing.T, from net.PacketConn, to net.PacketConn, msg []byte) {
|
||||
@@ -48,7 +46,6 @@ func TestStateUDPResponseReusesPriorCapturedValues(t *testing.T) {
|
||||
},
|
||||
},
|
||||
}
|
||||
maskManager := finalmask.NewUdpmaskManager([]finalmask.Udpmask{cfg})
|
||||
|
||||
clientRaw, err := net.ListenPacket("udp", "127.0.0.1:0")
|
||||
if err != nil {
|
||||
@@ -62,11 +59,11 @@ func TestStateUDPResponseReusesPriorCapturedValues(t *testing.T) {
|
||||
}
|
||||
defer serverRaw.Close()
|
||||
|
||||
client, err := maskManager.WrapPacketConnClient(clientRaw)
|
||||
client, err := cfg.WrapPacketConnClient(clientRaw, nil, nil)
|
||||
if err != nil {
|
||||
t.Fatal(err)
|
||||
}
|
||||
server, err := maskManager.WrapPacketConnServer(serverRaw)
|
||||
server, err := cfg.WrapPacketConnServer(serverRaw, nil, nil)
|
||||
if err != nil {
|
||||
t.Fatal(err)
|
||||
}
|
||||
|
||||
@@ -37,7 +37,7 @@ func TestDSLTCPHandshakeReusesCapturedValue(t *testing.T) {
|
||||
defer clientRaw.Close()
|
||||
defer serverRaw.Close()
|
||||
|
||||
client, err := cfg.WrapConnClient(clientRaw)
|
||||
client, err := cfg.WrapConnClient(clientRaw, nil, nil)
|
||||
if err != nil {
|
||||
t.Fatal(err)
|
||||
}
|
||||
@@ -117,7 +117,7 @@ func TestDSLTCPClientRejectsMismatchedResponseSequence(t *testing.T) {
|
||||
defer clientRaw.Close()
|
||||
defer serverRaw.Close()
|
||||
|
||||
client, err := clientCfg.WrapConnClient(clientRaw)
|
||||
client, err := clientCfg.WrapConnClient(clientRaw, nil, nil)
|
||||
if err != nil {
|
||||
t.Fatal(err)
|
||||
}
|
||||
|
||||
@@ -1,15 +1,16 @@
|
||||
package aes128gcm
|
||||
|
||||
import (
|
||||
"net"
|
||||
"github.com/xtls/xray-core/common/net"
|
||||
"github.com/xtls/xray-core/transport/internet/finalmask"
|
||||
)
|
||||
|
||||
func (c *Config) HeaderConn() {}
|
||||
|
||||
func (c *Config) WrapPacketConnClient(raw net.PacketConn, level int, levelCount int) (net.PacketConn, error) {
|
||||
return NewConnClient(c, raw)
|
||||
func (c *Config) WrapPacketConnClient(conn net.PacketConn, dest *net.Destination, dialer *finalmask.Dialer) (net.PacketConn, error) {
|
||||
return NewConnClient(c, conn)
|
||||
}
|
||||
|
||||
func (c *Config) WrapPacketConnServer(raw net.PacketConn, level int, levelCount int) (net.PacketConn, error) {
|
||||
return NewConnServer(c, raw)
|
||||
func (c *Config) WrapPacketConnServer(conn net.PacketConn, addr net.Addr, lc *finalmask.ListenConfig) (net.PacketConn, error) {
|
||||
return NewConnServer(c, conn)
|
||||
}
|
||||
|
||||
@@ -1,15 +1,16 @@
|
||||
package header
|
||||
|
||||
import (
|
||||
"net"
|
||||
"github.com/xtls/xray-core/common/net"
|
||||
"github.com/xtls/xray-core/transport/internet/finalmask"
|
||||
)
|
||||
|
||||
func (c *Config) HeaderConn() {}
|
||||
|
||||
func (c *Config) WrapPacketConnClient(raw net.PacketConn, level int, levelCount int) (net.PacketConn, error) {
|
||||
return NewConnClient(c, raw)
|
||||
func (c *Config) WrapPacketConnClient(conn net.PacketConn, dest *net.Destination, dialer *finalmask.Dialer) (net.PacketConn, error) {
|
||||
return NewConnClient(c, conn)
|
||||
}
|
||||
|
||||
func (c *Config) WrapPacketConnServer(raw net.PacketConn, level int, levelCount int) (net.PacketConn, error) {
|
||||
return NewConnServer(c, raw)
|
||||
func (c *Config) WrapPacketConnServer(conn net.PacketConn, addr net.Addr, lc *finalmask.ListenConfig) (net.PacketConn, error) {
|
||||
return NewConnServer(c, conn)
|
||||
}
|
||||
|
||||
@@ -1,15 +1,16 @@
|
||||
package original
|
||||
|
||||
import (
|
||||
"net"
|
||||
"github.com/xtls/xray-core/common/net"
|
||||
"github.com/xtls/xray-core/transport/internet/finalmask"
|
||||
)
|
||||
|
||||
func (c *Config) HeaderConn() {}
|
||||
|
||||
func (c *Config) WrapPacketConnClient(raw net.PacketConn, level int, levelCount int) (net.PacketConn, error) {
|
||||
return NewConnClient(c, raw)
|
||||
func (c *Config) WrapPacketConnClient(conn net.PacketConn, dest *net.Destination, dialer *finalmask.Dialer) (net.PacketConn, error) {
|
||||
return NewConnClient(c, conn)
|
||||
}
|
||||
|
||||
func (c *Config) WrapPacketConnServer(raw net.PacketConn, level int, levelCount int) (net.PacketConn, error) {
|
||||
return NewConnServer(c, raw)
|
||||
func (c *Config) WrapPacketConnServer(conn net.PacketConn, addr net.Addr, lc *finalmask.ListenConfig) (net.PacketConn, error) {
|
||||
return NewConnServer(c, conn)
|
||||
}
|
||||
|
||||
@@ -1,11 +1,14 @@
|
||||
package noise
|
||||
|
||||
import "net"
|
||||
import (
|
||||
"github.com/xtls/xray-core/common/net"
|
||||
"github.com/xtls/xray-core/transport/internet/finalmask"
|
||||
)
|
||||
|
||||
func (c *Config) WrapPacketConnClient(raw net.PacketConn, level int, levelCount int) (net.PacketConn, error) {
|
||||
return NewConnClient(c, raw)
|
||||
func (c *Config) WrapPacketConnClient(conn net.PacketConn, dest *net.Destination, dialer *finalmask.Dialer) (net.PacketConn, error) {
|
||||
return NewConnClient(c, conn)
|
||||
}
|
||||
|
||||
func (c *Config) WrapPacketConnServer(raw net.PacketConn, level int, levelCount int) (net.PacketConn, error) {
|
||||
return NewConnServer(c, raw)
|
||||
func (c *Config) WrapPacketConnServer(conn net.PacketConn, addr net.Addr, lc *finalmask.ListenConfig) (net.PacketConn, error) {
|
||||
return NewConnServer(c, conn)
|
||||
}
|
||||
|
||||
@@ -1,23 +1,14 @@
|
||||
package realm
|
||||
|
||||
import (
|
||||
"net"
|
||||
|
||||
"github.com/xtls/xray-core/common/errors"
|
||||
"github.com/xtls/xray-core/transport/internet"
|
||||
"github.com/xtls/xray-core/common/net"
|
||||
"github.com/xtls/xray-core/transport/internet/finalmask"
|
||||
)
|
||||
|
||||
func (c *Config) WrapPacketConnClient(raw net.PacketConn, level int, levelCount int) (net.PacketConn, error) {
|
||||
_, ok1 := raw.(*internet.FakePacketConn)
|
||||
if level != 0 || ok1 {
|
||||
return nil, errors.New("realm requires being at the outermost level")
|
||||
}
|
||||
return NewConnClient(c, raw)
|
||||
func (c *Config) WrapPacketConnClient(conn net.PacketConn, dest *net.Destination, dialer *finalmask.Dialer) (net.PacketConn, error) {
|
||||
return NewConnClient(c, conn)
|
||||
}
|
||||
|
||||
func (c *Config) WrapPacketConnServer(raw net.PacketConn, level int, levelCount int) (net.PacketConn, error) {
|
||||
if level != 0 {
|
||||
return nil, errors.New("realm requires being at the outermost level")
|
||||
}
|
||||
return NewConnServer(c, raw)
|
||||
func (c *Config) WrapPacketConnServer(conn net.PacketConn, addr net.Addr, lc *finalmask.ListenConfig) (net.PacketConn, error) {
|
||||
return NewConnServer(c, conn)
|
||||
}
|
||||
|
||||
@@ -1,23 +1,24 @@
|
||||
package salamander
|
||||
|
||||
import (
|
||||
"net"
|
||||
"github.com/xtls/xray-core/common/net"
|
||||
"github.com/xtls/xray-core/transport/internet/finalmask"
|
||||
)
|
||||
|
||||
func (c *Config) HeaderConn() {}
|
||||
|
||||
func (c *Config) WrapPacketConnClient(raw net.PacketConn, level int, levelCount int) (net.PacketConn, error) {
|
||||
return NewSalamanderConnClient(c, raw)
|
||||
func (c *Config) WrapPacketConnClient(conn net.PacketConn, dest *net.Destination, dialer *finalmask.Dialer) (net.PacketConn, error) {
|
||||
return NewSalamanderConnClient(c, conn)
|
||||
}
|
||||
|
||||
func (c *Config) WrapPacketConnServer(raw net.PacketConn, level int, levelCount int) (net.PacketConn, error) {
|
||||
return NewSalamanderConnServer(c, raw)
|
||||
func (c *Config) WrapPacketConnServer(conn net.PacketConn, addr net.Addr, lc *finalmask.ListenConfig) (net.PacketConn, error) {
|
||||
return NewSalamanderConnServer(c, conn)
|
||||
}
|
||||
|
||||
func (c *GeckoConfig) WrapPacketConnClient(raw net.PacketConn, level int, levelCount int) (net.PacketConn, error) {
|
||||
return NewGeckoConnClient(c, raw)
|
||||
func (c *GeckoConfig) WrapPacketConnClient(conn net.PacketConn, dest *net.Destination, dialer *finalmask.Dialer) (net.PacketConn, error) {
|
||||
return NewGeckoConnClient(c, conn)
|
||||
}
|
||||
|
||||
func (c *GeckoConfig) WrapPacketConnServer(raw net.PacketConn, level int, levelCount int) (net.PacketConn, error) {
|
||||
return NewGeckoConnServer(c, raw)
|
||||
func (c *GeckoConfig) WrapPacketConnServer(conn net.PacketConn, addr net.Addr, lc *finalmask.ListenConfig) (net.PacketConn, error) {
|
||||
return NewGeckoConnServer(c, conn)
|
||||
}
|
||||
|
||||
@@ -1,19 +1,18 @@
|
||||
package sudoku
|
||||
|
||||
import (
|
||||
"net"
|
||||
|
||||
"github.com/xtls/xray-core/common/errors"
|
||||
"github.com/xtls/xray-core/common/net"
|
||||
"github.com/xtls/xray-core/transport/internet/finalmask"
|
||||
)
|
||||
|
||||
// Sudoku in finalmask mode is a pure appearance transform with no standalone handshake.
|
||||
// TCP always keeps classic sudoku on uplink and uses packed downlink optimization on server writes.
|
||||
func (c *Config) WrapConnClient(raw net.Conn) (net.Conn, error) {
|
||||
return newPackedDirectionalConn(raw, c, true)
|
||||
func (c *Config) WrapConnClient(conn net.Conn, dest *net.Destination, dialer *finalmask.Dialer) (net.Conn, error) {
|
||||
return newPackedDirectionalConn(conn, c, true)
|
||||
}
|
||||
|
||||
func (c *Config) WrapConnServer(raw net.Conn) (net.Conn, error) {
|
||||
return newPackedDirectionalConn(raw, c, false)
|
||||
func (c *Config) WrapConnServer(conn net.Conn) (net.Conn, error) {
|
||||
return newPackedDirectionalConn(conn, c, false)
|
||||
}
|
||||
|
||||
func newPackedDirectionalConn(raw net.Conn, config *Config, readPacked bool) (net.Conn, error) {
|
||||
@@ -36,16 +35,10 @@ func newPackedDirectionalConn(raw net.Conn, config *Config, readPacked bool) (ne
|
||||
return newWrappedConn(raw, reader, writer), nil
|
||||
}
|
||||
|
||||
func (c *Config) WrapPacketConnClient(raw net.PacketConn, level int, levelCount int) (net.PacketConn, error) {
|
||||
if level != levelCount {
|
||||
return nil, errors.New("sudoku udp mask must be the innermost mask in chain")
|
||||
}
|
||||
return NewUDPConn(raw, c)
|
||||
func (c *Config) WrapPacketConnClient(conn net.PacketConn, dest *net.Destination, dialer *finalmask.Dialer) (net.PacketConn, error) {
|
||||
return NewUDPConn(conn, c)
|
||||
}
|
||||
|
||||
func (c *Config) WrapPacketConnServer(raw net.PacketConn, level int, levelCount int) (net.PacketConn, error) {
|
||||
if level != levelCount {
|
||||
return nil, errors.New("sudoku udp mask must be the innermost mask in chain")
|
||||
}
|
||||
return NewUDPConn(raw, c)
|
||||
func (c *Config) WrapPacketConnServer(conn net.PacketConn, addr net.Addr, lc *finalmask.ListenConfig) (net.PacketConn, error) {
|
||||
return NewUDPConn(conn, c)
|
||||
}
|
||||
|
||||
@@ -2,12 +2,14 @@ package finalmask_test
|
||||
|
||||
import (
|
||||
"bytes"
|
||||
"context"
|
||||
"io"
|
||||
"net"
|
||||
gonet "net"
|
||||
"strings"
|
||||
"testing"
|
||||
"time"
|
||||
|
||||
"github.com/xtls/xray-core/common/net"
|
||||
"github.com/xtls/xray-core/transport/internet/finalmask"
|
||||
"github.com/xtls/xray-core/transport/internet/finalmask/header/custom"
|
||||
)
|
||||
@@ -20,11 +22,14 @@ func mustSendRecvTcp(
|
||||
) {
|
||||
t.Helper()
|
||||
|
||||
waitCh := make(chan error)
|
||||
|
||||
go func() {
|
||||
_, err := from.Write(msg)
|
||||
if err != nil {
|
||||
t.Error(err)
|
||||
t.Fatal(err)
|
||||
}
|
||||
close(waitCh)
|
||||
}()
|
||||
|
||||
buf := make([]byte, 1024)
|
||||
@@ -40,18 +45,23 @@ func mustSendRecvTcp(
|
||||
if !bytes.Equal(buf[:n], msg) {
|
||||
t.Fatalf("unexpected data %q", buf[:n])
|
||||
}
|
||||
|
||||
<-waitCh
|
||||
}
|
||||
|
||||
type layerMaskTcp struct {
|
||||
name string
|
||||
mask finalmask.Tcpmask
|
||||
mask finalmask.TCPMask
|
||||
}
|
||||
|
||||
type failingWrapMask struct{}
|
||||
|
||||
func (failingWrapMask) TCP() {}
|
||||
func (f failingWrapMask) WrapConnClient(raw net.Conn) (net.Conn, error) { return raw, nil }
|
||||
func (f failingWrapMask) WrapConnServer(raw net.Conn) (net.Conn, error) {
|
||||
func (failingWrapMask) TCP() {}
|
||||
func (f failingWrapMask) WrapConnClient(conn net.Conn, dest *net.Destination, dialer *finalmask.Dialer) (net.Conn, error) {
|
||||
return conn, nil
|
||||
}
|
||||
|
||||
func (f failingWrapMask) WrapConnServer(conn net.Conn) (net.Conn, error) {
|
||||
return nil, io.ErrClosedPipe
|
||||
}
|
||||
|
||||
@@ -92,32 +102,31 @@ func TestConnReadWrite(t *testing.T) {
|
||||
t.Run(c.name, func(t *testing.T) {
|
||||
mask := c.mask
|
||||
|
||||
maskManager := finalmask.NewTcpmaskManager([]finalmask.Tcpmask{mask})
|
||||
dialTCP := func(ctx context.Context, dest net.Destination) (net.Conn, error) {
|
||||
return net.Dial("tcp", dest.NetAddr())
|
||||
}
|
||||
listen := func(ctx context.Context, addr net.Addr) (net.Listener, error) {
|
||||
return net.Listen("tcp", addr.String())
|
||||
}
|
||||
finalMask := finalmask.NewFinalMask([]finalmask.TCPMask{mask}, nil, dialTCP, listen, nil, nil)
|
||||
|
||||
ln, err := net.Listen("tcp", "127.0.0.1:0")
|
||||
listener, err := finalMask.Listen(context.Background(), &net.TCPAddr{IP: net.LocalHostIP.IP()})
|
||||
if err != nil {
|
||||
t.Fatal(err)
|
||||
}
|
||||
t.Cleanup(func() { listener.Close() })
|
||||
|
||||
client, err := net.Dial("tcp", ln.Addr().String())
|
||||
client, err := finalMask.DialTCP(context.Background(), net.TCPDestination(net.IPAddress(listener.Addr().(*net.TCPAddr).IP), net.Port(listener.Addr().(*net.TCPAddr).Port)))
|
||||
if err != nil {
|
||||
t.Fatal(err)
|
||||
}
|
||||
t.Cleanup(func() { client.Close() })
|
||||
|
||||
client, err = maskManager.WrapConnClient(client)
|
||||
if err != nil {
|
||||
t.Fatal(err)
|
||||
}
|
||||
|
||||
server, err := ln.Accept()
|
||||
if err != nil {
|
||||
t.Fatal(err)
|
||||
}
|
||||
|
||||
server, err = maskManager.WrapConnServer(server)
|
||||
server, err := listener.Accept()
|
||||
if err != nil {
|
||||
t.Fatal(err)
|
||||
}
|
||||
t.Cleanup(func() { server.Close() })
|
||||
|
||||
_ = client.SetDeadline(time.Now().Add(time.Second))
|
||||
_ = server.SetDeadline(time.Now().Add(time.Second))
|
||||
@@ -150,34 +159,32 @@ func TestTCPcustomStaticHandshakeRoundTrip(t *testing.T) {
|
||||
},
|
||||
},
|
||||
}
|
||||
maskManager := finalmask.NewTcpmaskManager([]finalmask.Tcpmask{cfg})
|
||||
|
||||
ln, err := net.Listen("tcp", "127.0.0.1:0")
|
||||
if err != nil {
|
||||
t.Fatal(err)
|
||||
dialTCP := func(ctx context.Context, dest net.Destination) (net.Conn, error) {
|
||||
return net.Dial("tcp", dest.NetAddr())
|
||||
}
|
||||
defer ln.Close()
|
||||
listen := func(ctx context.Context, addr net.Addr) (net.Listener, error) {
|
||||
return net.Listen("tcp", addr.String())
|
||||
}
|
||||
finalMask := finalmask.NewFinalMask([]finalmask.TCPMask{cfg}, nil, dialTCP, listen, nil, nil)
|
||||
|
||||
clientRaw, err := net.Dial("tcp", ln.Addr().String())
|
||||
listener, err := finalMask.Listen(context.Background(), &net.TCPAddr{IP: net.LocalHostIP.IP()})
|
||||
if err != nil {
|
||||
t.Fatal(err)
|
||||
}
|
||||
defer clientRaw.Close()
|
||||
defer listener.Close()
|
||||
|
||||
serverRaw, err := ln.Accept()
|
||||
client, err := finalMask.DialTCP(context.Background(), net.TCPDestination(net.IPAddress(listener.Addr().(*net.TCPAddr).IP), net.Port(listener.Addr().(*net.TCPAddr).Port)))
|
||||
if err != nil {
|
||||
t.Fatal(err)
|
||||
}
|
||||
defer serverRaw.Close()
|
||||
defer client.Close()
|
||||
|
||||
client, err := maskManager.WrapConnClient(clientRaw)
|
||||
if err != nil {
|
||||
t.Fatal(err)
|
||||
}
|
||||
server, err := maskManager.WrapConnServer(serverRaw)
|
||||
server, err := listener.Accept()
|
||||
if err != nil {
|
||||
t.Fatal(err)
|
||||
}
|
||||
defer server.Close()
|
||||
|
||||
_ = client.SetDeadline(time.Now().Add(time.Second))
|
||||
_ = server.SetDeadline(time.Now().Add(time.Second))
|
||||
@@ -220,11 +227,11 @@ func TestTCPcustomClientRejectsMismatchedServerSequence(t *testing.T) {
|
||||
},
|
||||
}
|
||||
|
||||
clientRaw, serverRaw := net.Pipe()
|
||||
clientRaw, serverRaw := gonet.Pipe()
|
||||
defer clientRaw.Close()
|
||||
defer serverRaw.Close()
|
||||
|
||||
client, err := clientCfg.WrapConnClient(clientRaw)
|
||||
client, err := clientCfg.WrapConnClient(clientRaw, nil, nil)
|
||||
if err != nil {
|
||||
t.Fatal(err)
|
||||
}
|
||||
@@ -257,42 +264,37 @@ func TestTCPcustomClientRejectsMismatchedServerSequence(t *testing.T) {
|
||||
}
|
||||
|
||||
func TestTCPWrapListenerRejectsImmediateWrapErrors(t *testing.T) {
|
||||
clientManager := finalmask.NewTcpmaskManager([]finalmask.Tcpmask{failingWrapMask{}})
|
||||
serverManager := finalmask.NewTcpmaskManager([]finalmask.Tcpmask{failingWrapMask{}})
|
||||
dialTCP := func(ctx context.Context, dest net.Destination) (net.Conn, error) {
|
||||
return net.Dial("tcp", dest.NetAddr())
|
||||
}
|
||||
listen := func(ctx context.Context, addr net.Addr) (net.Listener, error) {
|
||||
return net.Listen("tcp", addr.String())
|
||||
}
|
||||
finalMask := finalmask.NewFinalMask([]finalmask.TCPMask{failingWrapMask{}}, nil, dialTCP, listen, nil, nil)
|
||||
|
||||
rawLn, err := net.Listen("tcp", "127.0.0.1:0")
|
||||
if err != nil {
|
||||
t.Fatal(err)
|
||||
}
|
||||
defer rawLn.Close()
|
||||
|
||||
ln, err := serverManager.WrapListener(rawLn)
|
||||
listener, err := finalMask.Listen(context.Background(), &net.TCPAddr{IP: net.LocalHostIP.IP()})
|
||||
if err != nil {
|
||||
t.Fatal(err)
|
||||
}
|
||||
defer listener.Close()
|
||||
|
||||
accepted := make(chan struct {
|
||||
conn net.Conn
|
||||
err error
|
||||
}, 1)
|
||||
go func() {
|
||||
conn, err := ln.Accept()
|
||||
conn, err := listener.Accept()
|
||||
accepted <- struct {
|
||||
conn net.Conn
|
||||
err error
|
||||
}{conn: conn, err: err}
|
||||
}()
|
||||
|
||||
clientRaw, err := net.Dial("tcp", rawLn.Addr().String())
|
||||
if err != nil {
|
||||
t.Fatal(err)
|
||||
}
|
||||
defer clientRaw.Close()
|
||||
|
||||
client, err := clientManager.WrapConnClient(clientRaw)
|
||||
client, err := finalMask.DialTCP(context.Background(), net.TCPDestination(net.IPAddress(listener.Addr().(*net.TCPAddr).IP), net.Port(listener.Addr().(*net.TCPAddr).Port)))
|
||||
if err != nil {
|
||||
t.Fatal(err)
|
||||
}
|
||||
defer client.Close()
|
||||
|
||||
_ = client.SetDeadline(time.Now().Add(time.Second))
|
||||
|
||||
|
||||
@@ -2,13 +2,15 @@ package finalmask_test
|
||||
|
||||
import (
|
||||
"bytes"
|
||||
"context"
|
||||
"encoding/binary"
|
||||
"io"
|
||||
"net"
|
||||
gonet "net"
|
||||
"sync/atomic"
|
||||
"testing"
|
||||
"time"
|
||||
|
||||
"github.com/xtls/xray-core/common/net"
|
||||
"github.com/xtls/xray-core/proxy"
|
||||
"github.com/xtls/xray-core/transport/internet/finalmask"
|
||||
"github.com/xtls/xray-core/transport/internet/finalmask/header/custom"
|
||||
@@ -51,7 +53,7 @@ func mustSendRecv(
|
||||
|
||||
type layerMask struct {
|
||||
name string
|
||||
mask finalmask.Udpmask
|
||||
mask finalmask.UDPMask
|
||||
layers int
|
||||
}
|
||||
|
||||
@@ -213,25 +215,23 @@ func newStandaloneStunLikeUDPServerConfig() *custom.UDPStandaloneConfig {
|
||||
func newUDPClientServerPair(t *testing.T, cfg *custom.UDPStandaloneConfig) (net.PacketConn, net.PacketConn, net.PacketConn, net.PacketConn) {
|
||||
t.Helper()
|
||||
|
||||
clientRaw, err := net.ListenPacket("udp", "127.0.0.1:0")
|
||||
clientRaw, err := gonet.ListenPacket("udp", "127.0.0.1:0")
|
||||
if err != nil {
|
||||
t.Fatal(err)
|
||||
}
|
||||
t.Cleanup(func() { _ = clientRaw.Close() })
|
||||
|
||||
serverRaw, err := net.ListenPacket("udp", "127.0.0.1:0")
|
||||
serverRaw, err := gonet.ListenPacket("udp", "127.0.0.1:0")
|
||||
if err != nil {
|
||||
t.Fatal(err)
|
||||
}
|
||||
t.Cleanup(func() { _ = serverRaw.Close() })
|
||||
|
||||
maskManager := finalmask.NewUdpmaskManager([]finalmask.Udpmask{cfg})
|
||||
|
||||
client, err := maskManager.WrapPacketConnClient(clientRaw)
|
||||
client, err := cfg.WrapPacketConnClient(clientRaw, nil, nil)
|
||||
if err != nil {
|
||||
t.Fatal(err)
|
||||
}
|
||||
server, err := maskManager.WrapPacketConnServer(serverRaw)
|
||||
server, err := cfg.WrapPacketConnServer(serverRaw, nil, nil)
|
||||
if err != nil {
|
||||
t.Fatal(err)
|
||||
}
|
||||
@@ -348,31 +348,39 @@ func TestPacketConnReadWrite(t *testing.T) {
|
||||
if layers <= 0 {
|
||||
layers = 1
|
||||
}
|
||||
masks := make([]finalmask.Udpmask, 0, layers)
|
||||
masks := make([]finalmask.UDPMask, 0, layers)
|
||||
for i := 0; i < layers; i++ {
|
||||
masks = append(masks, mask)
|
||||
}
|
||||
maskManager := finalmask.NewUdpmaskManager(masks)
|
||||
|
||||
client, err := net.ListenPacket("udp", "127.0.0.1:0")
|
||||
dialUDP := func(ctx context.Context, dest net.Destination) (net.PacketConn, net.Addr, error) {
|
||||
udpAddr, err := net.ResolveUDPAddr("udp", dest.NetAddr())
|
||||
if err != nil {
|
||||
return nil, nil, err
|
||||
}
|
||||
conn, err := gonet.ListenPacket("udp", "127.0.0.1:0")
|
||||
if err != nil {
|
||||
return nil, nil, err
|
||||
}
|
||||
return conn, udpAddr, nil
|
||||
}
|
||||
listenPacket := func(ctx context.Context, addr net.Addr) (net.PacketConn, error) {
|
||||
return gonet.ListenPacket(addr.Network(), addr.String())
|
||||
}
|
||||
finalMask := finalmask.NewFinalMask(nil, masks, nil, nil, dialUDP, listenPacket)
|
||||
|
||||
server, err := finalMask.ListenPacket(context.Background(), &net.UDPAddr{IP: net.LocalHostIP.IP()})
|
||||
if err != nil {
|
||||
t.Fatal(err)
|
||||
}
|
||||
t.Cleanup(func() { server.Close() })
|
||||
|
||||
client, err = maskManager.WrapPacketConnClient(client)
|
||||
if err != nil {
|
||||
t.Fatal(err)
|
||||
}
|
||||
|
||||
server, err := net.ListenPacket("udp", "127.0.0.1:0")
|
||||
if err != nil {
|
||||
t.Fatal(err)
|
||||
}
|
||||
|
||||
server, err = maskManager.WrapPacketConnServer(server)
|
||||
clientConn, err := finalMask.DialUDP(context.Background(), net.UDPDestination(net.IPAddress(server.LocalAddr().(*net.UDPAddr).IP), net.Port(server.LocalAddr().(*net.UDPAddr).Port)))
|
||||
if err != nil {
|
||||
t.Fatal(err)
|
||||
}
|
||||
t.Cleanup(func() { clientConn.Close() })
|
||||
client := clientConn.(*finalmask.PacketConnWrapper).PacketConn
|
||||
|
||||
_ = client.SetDeadline(time.Now().Add(time.Second))
|
||||
_ = server.SetDeadline(time.Now().Add(time.Second))
|
||||
@@ -397,21 +405,20 @@ func TestUDPcustomStaticHeaderWireShape(t *testing.T) {
|
||||
{Rand: 1, RandMin: 0x30, RandMax: 0x40},
|
||||
},
|
||||
}
|
||||
maskManager := finalmask.NewUdpmaskManager([]finalmask.Udpmask{cfg})
|
||||
|
||||
clientRaw, err := net.ListenPacket("udp", "127.0.0.1:0")
|
||||
clientRaw, err := gonet.ListenPacket("udp", "127.0.0.1:0")
|
||||
if err != nil {
|
||||
t.Fatal(err)
|
||||
}
|
||||
defer clientRaw.Close()
|
||||
|
||||
serverRaw, err := net.ListenPacket("udp", "127.0.0.1:0")
|
||||
serverRaw, err := gonet.ListenPacket("udp", "127.0.0.1:0")
|
||||
if err != nil {
|
||||
t.Fatal(err)
|
||||
}
|
||||
defer serverRaw.Close()
|
||||
|
||||
client, err := maskManager.WrapPacketConnClient(clientRaw)
|
||||
client, err := cfg.WrapPacketConnClient(clientRaw, nil, nil)
|
||||
if err != nil {
|
||||
t.Fatal(err)
|
||||
}
|
||||
@@ -642,11 +649,11 @@ func TestSudokuBDD(t *testing.T) {
|
||||
Ascii: "prefer_ascii",
|
||||
}
|
||||
|
||||
clientRaw, serverRaw := net.Pipe()
|
||||
clientRaw, serverRaw := gonet.Pipe()
|
||||
defer clientRaw.Close()
|
||||
defer serverRaw.Close()
|
||||
|
||||
clientConn, err := cfg.WrapConnClient(clientRaw)
|
||||
clientConn, err := cfg.WrapConnClient(clientRaw, nil, nil)
|
||||
if err != nil {
|
||||
t.Fatal(err)
|
||||
}
|
||||
@@ -683,11 +690,11 @@ func TestSudokuBDD(t *testing.T) {
|
||||
PaddingMax: 0,
|
||||
}
|
||||
|
||||
clientRaw, serverRaw := net.Pipe()
|
||||
clientRaw, serverRaw := gonet.Pipe()
|
||||
defer clientRaw.Close()
|
||||
defer serverRaw.Close()
|
||||
|
||||
clientConn, err := cfg.WrapConnClient(clientRaw)
|
||||
clientConn, err := cfg.WrapConnClient(clientRaw, nil, nil)
|
||||
if err != nil {
|
||||
t.Fatal(err)
|
||||
}
|
||||
@@ -738,10 +745,10 @@ func TestSudokuBDD(t *testing.T) {
|
||||
countWireBytes := func(wrapServer func(net.Conn, *sudoku.Config) (net.Conn, error), cfg *sudoku.Config) int64 {
|
||||
t.Helper()
|
||||
|
||||
clientRaw, serverRaw := net.Pipe()
|
||||
clientRaw, serverRaw := gonet.Pipe()
|
||||
watchedServerRaw := &countingConn{Conn: serverRaw}
|
||||
|
||||
clientConn, err := cfg.WrapConnClient(clientRaw)
|
||||
clientConn, err := cfg.WrapConnClient(clientRaw, nil, nil)
|
||||
if err != nil {
|
||||
t.Fatal(err)
|
||||
}
|
||||
@@ -793,11 +800,11 @@ func TestSudokuBDD(t *testing.T) {
|
||||
CustomTables: []string{"xpxvvpvv", "vxpvxvvp"},
|
||||
}
|
||||
|
||||
clientRaw, serverRaw := net.Pipe()
|
||||
clientRaw, serverRaw := gonet.Pipe()
|
||||
defer clientRaw.Close()
|
||||
defer serverRaw.Close()
|
||||
|
||||
clientConn, err := cfg.WrapConnClient(clientRaw)
|
||||
clientConn, err := cfg.WrapConnClient(clientRaw, nil, nil)
|
||||
if err != nil {
|
||||
t.Fatal(err)
|
||||
}
|
||||
@@ -835,11 +842,11 @@ func TestSudokuBDD(t *testing.T) {
|
||||
PaddingMax: 0,
|
||||
}
|
||||
|
||||
clientRaw, serverRaw := net.Pipe()
|
||||
clientRaw, serverRaw := gonet.Pipe()
|
||||
defer clientRaw.Close()
|
||||
defer serverRaw.Close()
|
||||
|
||||
clientConn, err := cfg.WrapConnClient(clientRaw)
|
||||
clientConn, err := cfg.WrapConnClient(clientRaw, nil, nil)
|
||||
if err != nil {
|
||||
t.Fatal(err)
|
||||
}
|
||||
@@ -868,19 +875,6 @@ func TestSudokuBDD(t *testing.T) {
|
||||
}
|
||||
})
|
||||
|
||||
t.Run("GivenSudokuUDPMask_WhenNotInnermost_ThenWrapFails", func(t *testing.T) {
|
||||
cfg := &sudoku.Config{Password: "sudoku-udp"}
|
||||
raw, err := net.ListenPacket("udp", "127.0.0.1:0")
|
||||
if err != nil {
|
||||
t.Fatal(err)
|
||||
}
|
||||
defer raw.Close()
|
||||
|
||||
if _, err := cfg.WrapPacketConnClient(raw, 0, 1); err == nil {
|
||||
t.Fatal("expected innermost check failure")
|
||||
}
|
||||
})
|
||||
|
||||
t.Run("GivenSudokuMultiTableUDPMask_WhenClientSendsMultipleDatagrams_ThenPayloadMatches", func(t *testing.T) {
|
||||
cfg := &sudoku.Config{
|
||||
Password: "sudoku-udp-multi",
|
||||
@@ -889,25 +883,24 @@ func TestSudokuBDD(t *testing.T) {
|
||||
PaddingMin: 0,
|
||||
PaddingMax: 0,
|
||||
}
|
||||
maskManager := finalmask.NewUdpmaskManager([]finalmask.Udpmask{cfg})
|
||||
|
||||
clientRaw, err := net.ListenPacket("udp", "127.0.0.1:0")
|
||||
clientRaw, err := gonet.ListenPacket("udp", "127.0.0.1:0")
|
||||
if err != nil {
|
||||
t.Fatal(err)
|
||||
}
|
||||
defer clientRaw.Close()
|
||||
|
||||
serverRaw, err := net.ListenPacket("udp", "127.0.0.1:0")
|
||||
serverRaw, err := gonet.ListenPacket("udp", "127.0.0.1:0")
|
||||
if err != nil {
|
||||
t.Fatal(err)
|
||||
}
|
||||
defer serverRaw.Close()
|
||||
|
||||
client, err := maskManager.WrapPacketConnClient(clientRaw)
|
||||
client, err := cfg.WrapPacketConnClient(clientRaw, nil, nil)
|
||||
if err != nil {
|
||||
t.Fatal(err)
|
||||
}
|
||||
server, err := maskManager.WrapPacketConnServer(serverRaw)
|
||||
server, err := cfg.WrapPacketConnServer(serverRaw, nil, nil)
|
||||
if err != nil {
|
||||
t.Fatal(err)
|
||||
}
|
||||
@@ -961,7 +954,7 @@ func TestSudokuBDD(t *testing.T) {
|
||||
}
|
||||
defer serverRaw.Close()
|
||||
|
||||
clientConn, err := cfg.WrapConnClient(clientRaw)
|
||||
clientConn, err := cfg.WrapConnClient(clientRaw, nil, nil)
|
||||
if err != nil {
|
||||
t.Fatal(err)
|
||||
}
|
||||
@@ -1008,11 +1001,11 @@ func TestSudokuBDD(t *testing.T) {
|
||||
Ascii: "prefer_entropy",
|
||||
}
|
||||
|
||||
clientRaw, serverRaw := net.Pipe()
|
||||
clientRaw, serverRaw := gonet.Pipe()
|
||||
defer clientRaw.Close()
|
||||
defer serverRaw.Close()
|
||||
|
||||
clientConn, err := cfg.WrapConnClient(clientRaw)
|
||||
clientConn, err := cfg.WrapConnClient(clientRaw, nil, nil)
|
||||
if err != nil {
|
||||
t.Fatal(err)
|
||||
}
|
||||
@@ -1032,11 +1025,11 @@ func TestSudokuBDD(t *testing.T) {
|
||||
Ascii: "prefer_entropy",
|
||||
}
|
||||
|
||||
clientRaw, serverRaw := net.Pipe()
|
||||
clientRaw, serverRaw := gonet.Pipe()
|
||||
defer clientRaw.Close()
|
||||
defer serverRaw.Close()
|
||||
|
||||
clientConn, err := cfg.WrapConnClient(clientRaw)
|
||||
clientConn, err := cfg.WrapConnClient(clientRaw, nil, nil)
|
||||
if err != nil {
|
||||
t.Fatal(err)
|
||||
}
|
||||
|
||||
@@ -1,20 +1,17 @@
|
||||
package udphop
|
||||
|
||||
import (
|
||||
"net"
|
||||
|
||||
"github.com/xtls/xray-core/common/errors"
|
||||
"github.com/xtls/xray-core/transport/internet"
|
||||
"github.com/xtls/xray-core/common/net"
|
||||
"github.com/xtls/xray-core/transport/internet/finalmask"
|
||||
)
|
||||
|
||||
func (c *Config) WrapPacketConnClient(raw net.PacketConn, level int, levelCount int) (net.PacketConn, error) {
|
||||
_, ok1 := raw.(*internet.FakePacketConn)
|
||||
if level != 0 || ok1 {
|
||||
return nil, errors.New("udphop requires being at the outermost level")
|
||||
}
|
||||
return NewUDPHopConn(c, raw)
|
||||
func (c *Config) HandleDial() {}
|
||||
|
||||
func (c *Config) WrapPacketConnClient(conn net.PacketConn, dest *net.Destination, dialer *finalmask.Dialer) (net.PacketConn, error) {
|
||||
return NewUDPHopConn(c, dest, dialer)
|
||||
}
|
||||
|
||||
func (c *Config) WrapPacketConnServer(raw net.PacketConn, level int, levelCount int) (net.PacketConn, error) {
|
||||
func (c *Config) WrapPacketConnServer(conn net.PacketConn, addr net.Addr, lc *finalmask.ListenConfig) (net.PacketConn, error) {
|
||||
return nil, errors.New("udphop: client only")
|
||||
}
|
||||
|
||||
@@ -7,7 +7,6 @@
|
||||
package udphop
|
||||
|
||||
import (
|
||||
internet "github.com/xtls/xray-core/transport/internet"
|
||||
protoreflect "google.golang.org/protobuf/reflect/protoreflect"
|
||||
protoimpl "google.golang.org/protobuf/runtime/protoimpl"
|
||||
reflect "reflect"
|
||||
@@ -24,14 +23,13 @@ const (
|
||||
|
||||
type Config struct {
|
||||
state protoimpl.MessageState `protogen:"open.v1"`
|
||||
Sockopt *internet.SocketConfig `protobuf:"bytes,1,opt,name=sockopt,proto3" json:"sockopt,omitempty"`
|
||||
Local bool `protobuf:"varint,2,opt,name=local,proto3" json:"local,omitempty"`
|
||||
Remote bool `protobuf:"varint,3,opt,name=remote,proto3" json:"remote,omitempty"`
|
||||
RemoteOnce bool `protobuf:"varint,4,opt,name=remote_once,json=remoteOnce,proto3" json:"remote_once,omitempty"`
|
||||
IntervalMin int64 `protobuf:"varint,5,opt,name=interval_min,json=intervalMin,proto3" json:"interval_min,omitempty"`
|
||||
IntervalMax int64 `protobuf:"varint,6,opt,name=interval_max,json=intervalMax,proto3" json:"interval_max,omitempty"`
|
||||
RemotePorts []uint32 `protobuf:"varint,7,rep,packed,name=remote_ports,json=remotePorts,proto3" json:"remote_ports,omitempty"`
|
||||
RemoteIPs []string `protobuf:"bytes,8,rep,name=remoteIPs,proto3" json:"remoteIPs,omitempty"`
|
||||
RemoteIPs []string `protobuf:"bytes,7,rep,name=remoteIPs,proto3" json:"remoteIPs,omitempty"`
|
||||
RemotePorts []uint32 `protobuf:"varint,8,rep,packed,name=remote_ports,json=remotePorts,proto3" json:"remote_ports,omitempty"`
|
||||
unknownFields protoimpl.UnknownFields
|
||||
sizeCache protoimpl.SizeCache
|
||||
}
|
||||
@@ -66,13 +64,6 @@ func (*Config) Descriptor() ([]byte, []int) {
|
||||
return file_transport_internet_finalmask_udphop_config_proto_rawDescGZIP(), []int{0}
|
||||
}
|
||||
|
||||
func (x *Config) GetSockopt() *internet.SocketConfig {
|
||||
if x != nil {
|
||||
return x.Sockopt
|
||||
}
|
||||
return nil
|
||||
}
|
||||
|
||||
func (x *Config) GetLocal() bool {
|
||||
if x != nil {
|
||||
return x.Local
|
||||
@@ -108,16 +99,16 @@ func (x *Config) GetIntervalMax() int64 {
|
||||
return 0
|
||||
}
|
||||
|
||||
func (x *Config) GetRemotePorts() []uint32 {
|
||||
func (x *Config) GetRemoteIPs() []string {
|
||||
if x != nil {
|
||||
return x.RemotePorts
|
||||
return x.RemoteIPs
|
||||
}
|
||||
return nil
|
||||
}
|
||||
|
||||
func (x *Config) GetRemoteIPs() []string {
|
||||
func (x *Config) GetRemotePorts() []uint32 {
|
||||
if x != nil {
|
||||
return x.RemoteIPs
|
||||
return x.RemotePorts
|
||||
}
|
||||
return nil
|
||||
}
|
||||
@@ -126,17 +117,16 @@ var File_transport_internet_finalmask_udphop_config_proto protoreflect.FileDescr
|
||||
|
||||
const file_transport_internet_finalmask_udphop_config_proto_rawDesc = "" +
|
||||
"\n" +
|
||||
"0transport/internet/finalmask/udphop/config.proto\x12(xray.transport.internet.finalmask.udphop\x1a\x1ftransport/internet/config.proto\"\x9f\x02\n" +
|
||||
"\x06Config\x12?\n" +
|
||||
"\asockopt\x18\x01 \x01(\v2%.xray.transport.internet.SocketConfigR\asockopt\x12\x14\n" +
|
||||
"0transport/internet/finalmask/udphop/config.proto\x12(xray.transport.internet.finalmask.udphop\"\xe4\x01\n" +
|
||||
"\x06Config\x12\x14\n" +
|
||||
"\x05local\x18\x02 \x01(\bR\x05local\x12\x16\n" +
|
||||
"\x06remote\x18\x03 \x01(\bR\x06remote\x12\x1f\n" +
|
||||
"\vremote_once\x18\x04 \x01(\bR\n" +
|
||||
"remoteOnce\x12!\n" +
|
||||
"\finterval_min\x18\x05 \x01(\x03R\vintervalMin\x12!\n" +
|
||||
"\finterval_max\x18\x06 \x01(\x03R\vintervalMax\x12!\n" +
|
||||
"\fremote_ports\x18\a \x03(\rR\vremotePorts\x12\x1c\n" +
|
||||
"\tremoteIPs\x18\b \x03(\tR\tremoteIPsB\x9a\x01\n" +
|
||||
"\finterval_max\x18\x06 \x01(\x03R\vintervalMax\x12\x1c\n" +
|
||||
"\tremoteIPs\x18\a \x03(\tR\tremoteIPs\x12!\n" +
|
||||
"\fremote_ports\x18\b \x03(\rR\vremotePortsJ\x04\b\x01\x10\x02B\x9a\x01\n" +
|
||||
",com.xray.transport.internet.finalmask.udphopP\x01Z=github.com/xtls/xray-core/transport/internet/finalmask/udphop\xaa\x02(Xray.Transport.Internet.Finalmask.Udphopb\x06proto3"
|
||||
|
||||
var (
|
||||
@@ -153,16 +143,14 @@ func file_transport_internet_finalmask_udphop_config_proto_rawDescGZIP() []byte
|
||||
|
||||
var file_transport_internet_finalmask_udphop_config_proto_msgTypes = make([]protoimpl.MessageInfo, 1)
|
||||
var file_transport_internet_finalmask_udphop_config_proto_goTypes = []any{
|
||||
(*Config)(nil), // 0: xray.transport.internet.finalmask.udphop.Config
|
||||
(*internet.SocketConfig)(nil), // 1: xray.transport.internet.SocketConfig
|
||||
(*Config)(nil), // 0: xray.transport.internet.finalmask.udphop.Config
|
||||
}
|
||||
var file_transport_internet_finalmask_udphop_config_proto_depIdxs = []int32{
|
||||
1, // 0: xray.transport.internet.finalmask.udphop.Config.sockopt:type_name -> xray.transport.internet.SocketConfig
|
||||
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
|
||||
0, // [0:0] is the sub-list for method output_type
|
||||
0, // [0:0] is the sub-list for method input_type
|
||||
0, // [0:0] is the sub-list for extension type_name
|
||||
0, // [0:0] is the sub-list for extension extendee
|
||||
0, // [0:0] is the sub-list for field type_name
|
||||
}
|
||||
|
||||
func init() { file_transport_internet_finalmask_udphop_config_proto_init() }
|
||||
|
||||
@@ -6,16 +6,14 @@ option go_package = "github.com/xtls/xray-core/transport/internet/finalmask/udph
|
||||
option java_package = "com.xray.transport.internet.finalmask.udphop";
|
||||
option java_multiple_files = true;
|
||||
|
||||
import "transport/internet/config.proto";
|
||||
|
||||
message Config {
|
||||
xray.transport.internet.SocketConfig sockopt = 1;
|
||||
reserved 1;
|
||||
bool local = 2;
|
||||
bool remote = 3;
|
||||
bool remote_once = 4;
|
||||
int64 interval_min = 5;
|
||||
int64 interval_max = 6;
|
||||
repeated uint32 remote_ports = 7;
|
||||
repeated string remoteIPs = 8;
|
||||
repeated string remoteIPs = 7;
|
||||
repeated uint32 remote_ports = 8;
|
||||
}
|
||||
|
||||
|
||||
@@ -6,9 +6,7 @@ import (
|
||||
goerrors "errors"
|
||||
"io"
|
||||
mrand "math/rand"
|
||||
gonet "net"
|
||||
"net/netip"
|
||||
"reflect"
|
||||
"sync"
|
||||
"time"
|
||||
|
||||
@@ -16,8 +14,6 @@ import (
|
||||
"github.com/xtls/xray-core/common/crypto"
|
||||
"github.com/xtls/xray-core/common/errors"
|
||||
"github.com/xtls/xray-core/common/net"
|
||||
"github.com/xtls/xray-core/common/net/cnc"
|
||||
"github.com/xtls/xray-core/transport/internet"
|
||||
"github.com/xtls/xray-core/transport/internet/finalmask"
|
||||
)
|
||||
|
||||
@@ -34,16 +30,14 @@ type packet struct {
|
||||
}
|
||||
|
||||
type udpHopConn struct {
|
||||
conn net.PacketConn
|
||||
sockopt *internet.SocketConfig
|
||||
local bool
|
||||
remote bool
|
||||
remoteOnce bool
|
||||
dialer *finalmask.Dialer
|
||||
local bool
|
||||
remote bool
|
||||
|
||||
intervalMin int64
|
||||
intervalMax int64
|
||||
remotePorts []uint32
|
||||
remoteIPs []netip.Prefix
|
||||
remotePorts []uint32
|
||||
|
||||
deadline time.Time
|
||||
readDeadline time.Time
|
||||
@@ -55,10 +49,10 @@ type udpHopConn struct {
|
||||
readCh chan packet
|
||||
closeCh chan struct{}
|
||||
wg sync.WaitGroup
|
||||
mu sync.Mutex
|
||||
mu sync.RWMutex
|
||||
}
|
||||
|
||||
func NewUDPHopConn(c *Config, raw net.PacketConn) (net.PacketConn, error) {
|
||||
func NewUDPHopConn(c *Config, dest *net.Destination, dialer *finalmask.Dialer) (net.PacketConn, error) {
|
||||
if c.IntervalMin < 5 || c.IntervalMax < 5 {
|
||||
return nil, errors.New("invalid interval")
|
||||
}
|
||||
@@ -66,22 +60,40 @@ func NewUDPHopConn(c *Config, raw net.PacketConn) (net.PacketConn, error) {
|
||||
for _, ip := range c.RemoteIPs {
|
||||
remoteIPs = append(remoteIPs, netip.MustParsePrefix(ip))
|
||||
}
|
||||
conn := &udpHopConn{
|
||||
conn: raw,
|
||||
sockopt: c.Sockopt,
|
||||
local: c.Local,
|
||||
remote: c.Remote,
|
||||
remoteOnce: c.RemoteOnce,
|
||||
remotePorts := c.RemotePorts
|
||||
if c.Remote || c.RemoteOnce {
|
||||
if len(remoteIPs) > 0 {
|
||||
dest.Address = net.IPAddress(randPrefix(remoteIPs[mrand.Intn(len(remoteIPs))]))
|
||||
}
|
||||
if len(remotePorts) > 0 {
|
||||
dest.Port = net.Port(remotePorts[mrand.Intn(len(remotePorts))])
|
||||
}
|
||||
}
|
||||
conn, err := dialer.DialUDP(*dest)
|
||||
if err != nil {
|
||||
return nil, err
|
||||
}
|
||||
cur := conn.(*finalmask.PacketConnWrapper).PacketConn
|
||||
addr := conn.RemoteAddr().(*net.UDPAddr)
|
||||
client := &udpHopConn{
|
||||
dialer: dialer,
|
||||
local: c.Local,
|
||||
remote: c.Remote,
|
||||
|
||||
intervalMin: c.IntervalMin,
|
||||
intervalMax: c.IntervalMax,
|
||||
remotePorts: c.RemotePorts,
|
||||
remoteIPs: remoteIPs,
|
||||
remotePorts: remotePorts,
|
||||
|
||||
cur: cur,
|
||||
addr: addr,
|
||||
readCh: make(chan packet),
|
||||
closeCh: make(chan struct{}),
|
||||
}
|
||||
return conn, nil
|
||||
go client.run()
|
||||
client.wg.Add(1)
|
||||
go client.recv(client.cur)
|
||||
return client, nil
|
||||
}
|
||||
|
||||
func (c *udpHopConn) closed() bool {
|
||||
@@ -93,61 +105,67 @@ func (c *udpHopConn) closed() bool {
|
||||
}
|
||||
}
|
||||
|
||||
func (c *udpHopConn) hop(addr *net.UDPAddr) {
|
||||
func (c *udpHopConn) run() {
|
||||
ticker := time.NewTicker(time.Second * time.Duration(crypto.RandBetween(c.intervalMin, c.intervalMax+1)))
|
||||
defer ticker.Stop()
|
||||
for {
|
||||
select {
|
||||
case <-c.closeCh:
|
||||
return
|
||||
case <-ticker.C:
|
||||
ticker.Reset(time.Second * time.Duration(crypto.RandBetween(c.intervalMin, c.intervalMax+1)))
|
||||
c.hop()
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
func (c *udpHopConn) hop() {
|
||||
c.mu.Lock()
|
||||
defer c.mu.Unlock()
|
||||
if c.closed() {
|
||||
return
|
||||
}
|
||||
newAddr := &net.UDPAddr{IP: addr.IP, Port: addr.Port}
|
||||
newConn := c.conn
|
||||
if c.remote || c.remoteOnce && c.addr == nil {
|
||||
if len(c.remotePorts) > 0 {
|
||||
newAddr.Port = int(c.remotePorts[mrand.Intn(len(c.remotePorts))])
|
||||
}
|
||||
oldIP := c.addr.IP
|
||||
oldPort := c.addr.Port
|
||||
if c.remote {
|
||||
if len(c.remoteIPs) > 0 {
|
||||
newAddr.IP = randPrefix(c.remoteIPs[mrand.Intn(len(c.remoteIPs))])
|
||||
c.addr.IP = randPrefix(c.remoteIPs[mrand.Intn(len(c.remoteIPs))])
|
||||
}
|
||||
if len(c.remotePorts) > 0 {
|
||||
c.addr.Port = int(c.remotePorts[mrand.Intn(len(c.remotePorts))])
|
||||
}
|
||||
}
|
||||
if c.local {
|
||||
raw, err := internet.DialSystem(context.Background(), net.UDPDestination(net.IPAddress(newAddr.IP), net.Port(newAddr.Port)), c.sockopt)
|
||||
conn, err := c.dialer.DialUDP(net.UDPDestination(net.IPAddress(c.addr.IP), net.Port(c.addr.Port)))
|
||||
if err != nil {
|
||||
c.addr.IP = oldIP
|
||||
c.addr.Port = oldPort
|
||||
errors.LogErrorInner(context.Background(), err, "hop err")
|
||||
return
|
||||
}
|
||||
switch c := raw.(type) {
|
||||
case *internet.PacketConnWrapper:
|
||||
newConn = c.PacketConn
|
||||
case *cnc.Connection:
|
||||
newConn = &internet.FakePacketConn{Conn: c}
|
||||
default:
|
||||
panic(reflect.TypeOf(c))
|
||||
}
|
||||
newConn.SetDeadline(c.deadline)
|
||||
newConn.SetReadDeadline(c.readDeadline)
|
||||
newConn.SetWriteDeadline(c.writeDeadline)
|
||||
conn.SetDeadline(c.deadline)
|
||||
conn.SetReadDeadline(c.readDeadline)
|
||||
conn.SetWriteDeadline(c.writeDeadline)
|
||||
if c.pre != nil {
|
||||
_ = c.pre.Close()
|
||||
}
|
||||
c.pre = c.cur
|
||||
c.cur = conn.(*finalmask.PacketConnWrapper).PacketConn
|
||||
c.wg.Add(1)
|
||||
go c.recv(newConn)
|
||||
go c.recv(c.cur)
|
||||
}
|
||||
c.addr = newAddr
|
||||
c.cur = newConn
|
||||
}
|
||||
|
||||
func (c *udpHopConn) recv(conn net.PacketConn) {
|
||||
defer c.wg.Done()
|
||||
|
||||
for {
|
||||
if c.closed() {
|
||||
return
|
||||
}
|
||||
p := pool.Get().([]byte)
|
||||
n, addr, err := conn.ReadFrom(p)
|
||||
if err != nil {
|
||||
pool.Put(p[:cap(p)])
|
||||
if goerrors.Is(err, io.EOF) || goerrors.Is(err, io.ErrClosedPipe) || goerrors.Is(err, gonet.ErrClosed) {
|
||||
break
|
||||
if c.closed() {
|
||||
return
|
||||
}
|
||||
var netErr net.Error
|
||||
if goerrors.As(err, &netErr) && netErr.Timeout() {
|
||||
@@ -156,9 +174,10 @@ func (c *udpHopConn) recv(conn net.PacketConn) {
|
||||
case <-c.closeCh:
|
||||
return
|
||||
}
|
||||
continue
|
||||
}
|
||||
errors.LogErrorInner(context.Background(), err, "recv err")
|
||||
continue
|
||||
return
|
||||
}
|
||||
select {
|
||||
case c.readCh <- packet{p: p[:n], addr: addr}:
|
||||
@@ -169,22 +188,6 @@ func (c *udpHopConn) recv(conn net.PacketConn) {
|
||||
}
|
||||
}
|
||||
|
||||
func (c *udpHopConn) hopLoop() {
|
||||
ticker := time.NewTicker(time.Second * time.Duration(crypto.RandBetween(c.intervalMin, c.intervalMax+1)))
|
||||
defer ticker.Stop()
|
||||
for {
|
||||
select {
|
||||
case <-ticker.C:
|
||||
ticker.Reset(time.Second * time.Duration(crypto.RandBetween(c.intervalMin, c.intervalMax+1)))
|
||||
c.mu.Lock()
|
||||
c.hop(c.addr)
|
||||
c.mu.Unlock()
|
||||
case <-c.closeCh:
|
||||
return
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
func (c *udpHopConn) ReadFrom(p []byte) (n int, addr net.Addr, err error) {
|
||||
packet, ok := <-c.readCh
|
||||
if ok {
|
||||
@@ -194,21 +197,12 @@ func (c *udpHopConn) ReadFrom(p []byte) (n int, addr net.Addr, err error) {
|
||||
}
|
||||
return n, packet.addr, packet.err
|
||||
}
|
||||
return 0, nil, io.EOF
|
||||
return 0, nil, io.ErrClosedPipe
|
||||
}
|
||||
|
||||
func (c *udpHopConn) WriteTo(p []byte, addr net.Addr) (n int, err error) {
|
||||
c.mu.Lock()
|
||||
defer c.mu.Unlock()
|
||||
|
||||
if c.cur == nil {
|
||||
c.hop(addr.(*net.UDPAddr))
|
||||
if c.cur == nil {
|
||||
return 0, nil
|
||||
}
|
||||
go c.hopLoop()
|
||||
}
|
||||
|
||||
c.mu.RLock()
|
||||
defer c.mu.RUnlock()
|
||||
_, err = c.cur.WriteTo(p, c.addr)
|
||||
if err != nil {
|
||||
errors.LogErrorInner(context.Background(), err, "send err")
|
||||
@@ -227,15 +221,12 @@ func (c *udpHopConn) Close() error {
|
||||
if c.pre != nil {
|
||||
_ = c.pre.Close()
|
||||
}
|
||||
if c.cur != nil {
|
||||
_ = c.cur.Close()
|
||||
}
|
||||
_ = c.conn.Close()
|
||||
_ = c.cur.Close()
|
||||
c.wg.Wait()
|
||||
select {
|
||||
case p := <-c.readCh:
|
||||
if p.p != nil {
|
||||
pool.Put(p.p[:cap(p.p)])
|
||||
case packet := <-c.readCh:
|
||||
if packet.p != nil {
|
||||
pool.Put(packet.p[:cap(packet.p)])
|
||||
}
|
||||
default:
|
||||
}
|
||||
@@ -244,7 +235,9 @@ func (c *udpHopConn) Close() error {
|
||||
}
|
||||
|
||||
func (c *udpHopConn) LocalAddr() net.Addr {
|
||||
return c.conn.LocalAddr()
|
||||
c.mu.RLock()
|
||||
defer c.mu.RUnlock()
|
||||
return c.cur.LocalAddr()
|
||||
}
|
||||
|
||||
func (c *udpHopConn) SetDeadline(t time.Time) error {
|
||||
@@ -254,10 +247,7 @@ func (c *udpHopConn) SetDeadline(t time.Time) error {
|
||||
if c.pre != nil {
|
||||
_ = c.pre.SetDeadline(t)
|
||||
}
|
||||
if c.cur != nil {
|
||||
_ = c.cur.SetDeadline(t)
|
||||
}
|
||||
return nil
|
||||
return c.cur.SetDeadline(t)
|
||||
}
|
||||
|
||||
func (c *udpHopConn) SetReadDeadline(t time.Time) error {
|
||||
@@ -267,10 +257,7 @@ func (c *udpHopConn) SetReadDeadline(t time.Time) error {
|
||||
if c.pre != nil {
|
||||
_ = c.pre.SetReadDeadline(t)
|
||||
}
|
||||
if c.cur != nil {
|
||||
_ = c.cur.SetReadDeadline(t)
|
||||
}
|
||||
return nil
|
||||
return c.cur.SetReadDeadline(t)
|
||||
}
|
||||
|
||||
func (c *udpHopConn) SetWriteDeadline(t time.Time) error {
|
||||
@@ -280,10 +267,7 @@ func (c *udpHopConn) SetWriteDeadline(t time.Time) error {
|
||||
if c.pre != nil {
|
||||
_ = c.pre.SetWriteDeadline(t)
|
||||
}
|
||||
if c.cur != nil {
|
||||
_ = c.cur.SetWriteDeadline(t)
|
||||
}
|
||||
return nil
|
||||
return c.cur.SetWriteDeadline(t)
|
||||
}
|
||||
|
||||
func randPrefix(p netip.Prefix) []byte {
|
||||
|
||||
@@ -1,21 +1,14 @@
|
||||
package xdns
|
||||
|
||||
import (
|
||||
"net"
|
||||
"github.com/xtls/xray-core/common/net"
|
||||
"github.com/xtls/xray-core/transport/internet/finalmask"
|
||||
)
|
||||
|
||||
func (c *Config) WrapPacketConnClient(raw net.PacketConn, level int, levelCount int) (net.PacketConn, error) {
|
||||
// _, ok1 := raw.(*internet.FakePacketConn)
|
||||
// _, ok2 := raw.(*udphop.UdpHopPacketConn)
|
||||
// if level != 0 || ok1 || ok2 {
|
||||
// return nil, errors.New("xdns requires being at the outermost level")
|
||||
// }
|
||||
return NewConnClient(c, raw)
|
||||
func (c *Config) WrapPacketConnClient(conn net.PacketConn, dest *net.Destination, dialer *finalmask.Dialer) (net.PacketConn, error) {
|
||||
return NewConnClient(c, conn)
|
||||
}
|
||||
|
||||
func (c *Config) WrapPacketConnServer(raw net.PacketConn, level int, levelCount int) (net.PacketConn, error) {
|
||||
// if level != 0 {
|
||||
// return nil, errors.New("xdns requires being at the outermost level")
|
||||
// }
|
||||
return NewConnServer(c, raw)
|
||||
func (c *Config) WrapPacketConnServer(conn net.PacketConn, addr net.Addr, lc *finalmask.ListenConfig) (net.PacketConn, error) {
|
||||
return NewConnServer(c, conn)
|
||||
}
|
||||
|
||||
@@ -8,8 +8,7 @@ import (
|
||||
goerrors "errors"
|
||||
"fmt"
|
||||
"io"
|
||||
mathrand "math/rand"
|
||||
"net"
|
||||
mrand "math/rand"
|
||||
"net/netip"
|
||||
"sync"
|
||||
"time"
|
||||
@@ -17,6 +16,7 @@ import (
|
||||
|
||||
"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/finalmask"
|
||||
"golang.org/x/net/icmp"
|
||||
"golang.org/x/net/ipv4"
|
||||
@@ -36,11 +36,11 @@ type packet struct {
|
||||
}
|
||||
|
||||
type xicmpConnClient struct {
|
||||
conn net.PacketConn
|
||||
icmp4 *icmp.PacketConn
|
||||
icmp6 *icmp.PacketConn
|
||||
udp bool
|
||||
ips []netip.Addr
|
||||
ip net.IP
|
||||
clientID [8]byte
|
||||
id int
|
||||
seq int
|
||||
@@ -50,7 +50,7 @@ type xicmpConnClient struct {
|
||||
mu sync.Mutex
|
||||
}
|
||||
|
||||
func NewConnClient(c *Config, raw net.PacketConn) (net.PacketConn, error) {
|
||||
func NewConnClient(c *Config, dest *net.Destination) (net.PacketConn, error) {
|
||||
var icmp4, icmp6 *icmp.PacketConn
|
||||
var err4, err6 error
|
||||
if c.DGRAM {
|
||||
@@ -69,17 +69,24 @@ func NewConnClient(c *Config, raw net.PacketConn) (net.PacketConn, error) {
|
||||
ips = append(ips, netip.MustParseAddr(ip))
|
||||
}
|
||||
|
||||
var ip net.IP
|
||||
if len(ips) > 0 {
|
||||
ip = ips[mrand.Intn(len(ips))].AsSlice()
|
||||
} else {
|
||||
ip = dest.Address.IP()
|
||||
}
|
||||
|
||||
var clientID [8]byte
|
||||
common.Must2(rand.Read(clientID[:]))
|
||||
|
||||
conn := &xicmpConnClient{
|
||||
conn: raw,
|
||||
icmp4: icmp4,
|
||||
icmp6: icmp6,
|
||||
udp: c.DGRAM,
|
||||
ips: ips,
|
||||
ip: ip,
|
||||
clientID: clientID,
|
||||
id: mathrand.Intn(65536),
|
||||
id: mrand.Intn(65536),
|
||||
seq: 1,
|
||||
readCh: make(chan packet),
|
||||
closeCh: make(chan struct{}),
|
||||
@@ -92,10 +99,6 @@ func NewConnClient(c *Config, raw net.PacketConn) (net.PacketConn, error) {
|
||||
return conn, nil
|
||||
}
|
||||
|
||||
func (c *xicmpConnClient) ring(a, b uint16) uint16 {
|
||||
return min(a-b, b-a)
|
||||
}
|
||||
|
||||
func (c *xicmpConnClient) closed() bool {
|
||||
select {
|
||||
case <-c.closeCh:
|
||||
@@ -110,12 +113,11 @@ func (c *xicmpConnClient) recv4() {
|
||||
|
||||
var b [finalmask.UDPSize]byte
|
||||
for {
|
||||
if c.closed() {
|
||||
return
|
||||
}
|
||||
|
||||
n, addr, err := c.icmp4.ReadFrom(b[:])
|
||||
if err != nil {
|
||||
if c.closed() {
|
||||
return
|
||||
}
|
||||
var netErr net.Error
|
||||
if goerrors.As(err, &netErr) && netErr.Timeout() {
|
||||
select {
|
||||
@@ -125,9 +127,10 @@ func (c *xicmpConnClient) recv4() {
|
||||
case <-c.closeCh:
|
||||
return
|
||||
}
|
||||
continue
|
||||
}
|
||||
errors.LogErrorInner(context.Background(), err, "recv4 err")
|
||||
continue
|
||||
errors.LogErrorInner(context.Background(), err, "recv err 4")
|
||||
return
|
||||
}
|
||||
|
||||
msg, err := icmp.ParseMessage(1, b[:n])
|
||||
@@ -150,10 +153,6 @@ func (c *xicmpConnClient) recv4() {
|
||||
continue
|
||||
}
|
||||
|
||||
if c.ring(uint16(echo.Seq), uint16(c.seq)) > 1000 {
|
||||
continue
|
||||
}
|
||||
|
||||
if len(echo.Data) > 8 && bytes.Equal(echo.Data[:8], c.clientID[:]) {
|
||||
continue
|
||||
}
|
||||
@@ -182,12 +181,11 @@ func (c *xicmpConnClient) recv6() {
|
||||
|
||||
var b [finalmask.UDPSize]byte
|
||||
for {
|
||||
if c.closed() {
|
||||
return
|
||||
}
|
||||
|
||||
n, addr, err := c.icmp6.ReadFrom(b[:])
|
||||
if err != nil {
|
||||
if c.closed() {
|
||||
return
|
||||
}
|
||||
var netErr net.Error
|
||||
if goerrors.As(err, &netErr) && netErr.Timeout() {
|
||||
select {
|
||||
@@ -197,9 +195,10 @@ func (c *xicmpConnClient) recv6() {
|
||||
case <-c.closeCh:
|
||||
return
|
||||
}
|
||||
continue
|
||||
}
|
||||
errors.LogErrorInner(context.Background(), err, "recv6 err")
|
||||
continue
|
||||
errors.LogErrorInner(context.Background(), err, "recv err 6")
|
||||
return
|
||||
}
|
||||
|
||||
msg, err := icmp.ParseMessage(58, b[:n])
|
||||
@@ -222,10 +221,6 @@ func (c *xicmpConnClient) recv6() {
|
||||
continue
|
||||
}
|
||||
|
||||
if c.ring(uint16(echo.Seq), uint16(c.seq)) > 1000 {
|
||||
continue
|
||||
}
|
||||
|
||||
if len(echo.Data) > 8 && bytes.Equal(echo.Data[:8], c.clientID[:]) {
|
||||
continue
|
||||
}
|
||||
@@ -273,9 +268,9 @@ func (c *xicmpConnClient) WriteTo(p []byte, addr net.Addr) (n int, err error) {
|
||||
c.seq %= 65536
|
||||
c.mu.Unlock()
|
||||
|
||||
ip := addr.(*net.UDPAddr).IP
|
||||
ip := c.ip
|
||||
if len(c.ips) > 0 {
|
||||
ip = c.ips[mathrand.Intn(len(c.ips))].AsSlice()
|
||||
ip = c.ips[mrand.Intn(len(c.ips))].AsSlice()
|
||||
}
|
||||
|
||||
if c.udp {
|
||||
@@ -314,7 +309,6 @@ func (c *xicmpConnClient) Close() error {
|
||||
close(c.closeCh)
|
||||
_ = c.icmp4.Close()
|
||||
_ = c.icmp6.Close()
|
||||
_ = c.conn.Close()
|
||||
c.wg.Wait()
|
||||
select {
|
||||
case p := <-c.readCh:
|
||||
@@ -328,7 +322,7 @@ func (c *xicmpConnClient) Close() error {
|
||||
}
|
||||
|
||||
func (c *xicmpConnClient) LocalAddr() net.Addr {
|
||||
return c.conn.LocalAddr()
|
||||
return &net.UDPAddr{IP: []byte{0, 0, 0, 0}}
|
||||
}
|
||||
|
||||
func (c *xicmpConnClient) SetDeadline(t time.Time) error {
|
||||
|
||||
@@ -1,23 +1,23 @@
|
||||
package xicmp
|
||||
|
||||
import (
|
||||
"net"
|
||||
"errors"
|
||||
|
||||
"github.com/xtls/xray-core/common/errors"
|
||||
"github.com/xtls/xray-core/transport/internet"
|
||||
"github.com/xtls/xray-core/common/net"
|
||||
"github.com/xtls/xray-core/transport/internet/finalmask"
|
||||
)
|
||||
|
||||
func (c *Config) WrapPacketConnClient(raw net.PacketConn, level int, levelCount int) (net.PacketConn, error) {
|
||||
_, ok1 := raw.(*internet.FakePacketConn)
|
||||
if level != 0 || ok1 {
|
||||
return nil, errors.New("xicmp requires being at the outermost level")
|
||||
func (c *Config) HandleDial() {}
|
||||
|
||||
func (c *Config) HandleListen() {}
|
||||
|
||||
func (c *Config) WrapPacketConnClient(conn net.PacketConn, dest *net.Destination, dialer *finalmask.Dialer) (net.PacketConn, error) {
|
||||
if dest.Address.Family().IsDomain() && len(c.IPs) == 0 {
|
||||
return nil, errors.New("empty ip addresses")
|
||||
}
|
||||
return NewConnClient(c, raw)
|
||||
return NewConnClient(c, dest)
|
||||
}
|
||||
|
||||
func (c *Config) WrapPacketConnServer(raw net.PacketConn, level int, levelCount int) (net.PacketConn, error) {
|
||||
if level != 0 {
|
||||
return nil, errors.New("xicmp requires being at the outermost level")
|
||||
}
|
||||
return NewConnServer(c, raw)
|
||||
func (c *Config) WrapPacketConnServer(conn net.PacketConn, addr net.Addr, lc *finalmask.ListenConfig) (net.PacketConn, error) {
|
||||
return NewConnServer(c)
|
||||
}
|
||||
|
||||
@@ -37,7 +37,6 @@ type record struct {
|
||||
}
|
||||
|
||||
type xicmpConnServer struct {
|
||||
conn net.PacketConn
|
||||
icmp4 *icmp.PacketConn
|
||||
icmp6 *icmp.PacketConn
|
||||
ips map[netip.Addr]struct{}
|
||||
@@ -48,7 +47,7 @@ type xicmpConnServer struct {
|
||||
mu sync.Mutex
|
||||
}
|
||||
|
||||
func NewConnServer(c *Config, raw net.PacketConn) (net.PacketConn, error) {
|
||||
func NewConnServer(c *Config) (net.PacketConn, error) {
|
||||
icmp4, err := icmp.ListenPacket("ip4:icmp", "0.0.0.0")
|
||||
if err != nil {
|
||||
return nil, err
|
||||
@@ -64,7 +63,6 @@ func NewConnServer(c *Config, raw net.PacketConn) (net.PacketConn, error) {
|
||||
}
|
||||
|
||||
conn := &xicmpConnServer{
|
||||
conn: raw,
|
||||
icmp4: icmp4,
|
||||
icmp6: icmp6,
|
||||
ips: ips,
|
||||
@@ -115,12 +113,11 @@ func (c *xicmpConnServer) recv4() {
|
||||
|
||||
var b [finalmask.UDPSize]byte
|
||||
for {
|
||||
if c.closed() {
|
||||
return
|
||||
}
|
||||
|
||||
n, addr, err := c.icmp4.ReadFrom(b[:])
|
||||
if err != nil {
|
||||
if c.closed() {
|
||||
return
|
||||
}
|
||||
var netErr net.Error
|
||||
if goerrors.As(err, &netErr) && netErr.Timeout() {
|
||||
select {
|
||||
@@ -130,9 +127,10 @@ func (c *xicmpConnServer) recv4() {
|
||||
case <-c.closeCh:
|
||||
return
|
||||
}
|
||||
continue
|
||||
}
|
||||
errors.LogErrorInner(context.Background(), err, "recv4 err")
|
||||
continue
|
||||
errors.LogErrorInner(context.Background(), err, "recv err 4")
|
||||
return
|
||||
}
|
||||
|
||||
msg, err := icmp.ParseMessage(1, b[:n])
|
||||
@@ -195,12 +193,11 @@ func (c *xicmpConnServer) recv6() {
|
||||
|
||||
var b [finalmask.UDPSize]byte
|
||||
for {
|
||||
if c.closed() {
|
||||
return
|
||||
}
|
||||
|
||||
n, addr, err := c.icmp6.ReadFrom(b[:])
|
||||
if err != nil {
|
||||
if c.closed() {
|
||||
return
|
||||
}
|
||||
var netErr net.Error
|
||||
if goerrors.As(err, &netErr) && netErr.Timeout() {
|
||||
select {
|
||||
@@ -210,9 +207,10 @@ func (c *xicmpConnServer) recv6() {
|
||||
case <-c.closeCh:
|
||||
return
|
||||
}
|
||||
continue
|
||||
}
|
||||
errors.LogErrorInner(context.Background(), err, "recv6 err")
|
||||
continue
|
||||
errors.LogErrorInner(context.Background(), err, "recv err 6")
|
||||
return
|
||||
}
|
||||
|
||||
msg, err := icmp.ParseMessage(58, b[:n])
|
||||
@@ -330,7 +328,6 @@ func (c *xicmpConnServer) Close() error {
|
||||
close(c.closeCh)
|
||||
_ = c.icmp4.Close()
|
||||
_ = c.icmp6.Close()
|
||||
_ = c.conn.Close()
|
||||
c.wg.Wait()
|
||||
select {
|
||||
case p := <-c.readCh:
|
||||
@@ -344,7 +341,7 @@ func (c *xicmpConnServer) Close() error {
|
||||
}
|
||||
|
||||
func (c *xicmpConnServer) LocalAddr() net.Addr {
|
||||
return c.conn.LocalAddr()
|
||||
return &net.UDPAddr{IP: []byte{0, 0, 0, 0}}
|
||||
}
|
||||
|
||||
func (c *xicmpConnServer) SetDeadline(t time.Time) error {
|
||||
|
||||
@@ -39,7 +39,6 @@ type record struct {
|
||||
}
|
||||
|
||||
type xicmpConnServer struct {
|
||||
conn net.PacketConn
|
||||
icmp4 *icmp.PacketConn
|
||||
icmp6 *icmp.PacketConn
|
||||
ipv4PC *ipv4.PacketConn
|
||||
@@ -52,7 +51,7 @@ type xicmpConnServer struct {
|
||||
mu sync.Mutex
|
||||
}
|
||||
|
||||
func NewConnServer(c *Config, raw net.PacketConn) (net.PacketConn, error) {
|
||||
func NewConnServer(c *Config) (net.PacketConn, error) {
|
||||
icmp4, err := icmp.ListenPacket("ip4:icmp", "0.0.0.0")
|
||||
if err != nil {
|
||||
return nil, err
|
||||
@@ -68,7 +67,6 @@ func NewConnServer(c *Config, raw net.PacketConn) (net.PacketConn, error) {
|
||||
}
|
||||
|
||||
conn := &xicmpConnServer{
|
||||
conn: raw,
|
||||
icmp4: icmp4,
|
||||
icmp6: icmp6,
|
||||
ipv4PC: icmp4.IPv4PacketConn(),
|
||||
@@ -124,12 +122,11 @@ func (c *xicmpConnServer) recv4() {
|
||||
|
||||
var b [finalmask.UDPSize]byte
|
||||
for {
|
||||
if c.closed() {
|
||||
return
|
||||
}
|
||||
|
||||
n, cm, addr, err := c.ipv4PC.ReadFrom(b[:])
|
||||
if err != nil {
|
||||
if c.closed() {
|
||||
return
|
||||
}
|
||||
var netErr net.Error
|
||||
if goerrors.As(err, &netErr) && netErr.Timeout() {
|
||||
select {
|
||||
@@ -139,9 +136,10 @@ func (c *xicmpConnServer) recv4() {
|
||||
case <-c.closeCh:
|
||||
return
|
||||
}
|
||||
continue
|
||||
}
|
||||
errors.LogErrorInner(context.Background(), err, "recv4 err")
|
||||
continue
|
||||
errors.LogErrorInner(context.Background(), err, "recv err 4")
|
||||
return
|
||||
}
|
||||
|
||||
msg, err := icmp.ParseMessage(1, b[:n])
|
||||
@@ -205,12 +203,11 @@ func (c *xicmpConnServer) recv6() {
|
||||
|
||||
var b [finalmask.UDPSize]byte
|
||||
for {
|
||||
if c.closed() {
|
||||
return
|
||||
}
|
||||
|
||||
n, cm, addr, err := c.ipv6PC.ReadFrom(b[:])
|
||||
if err != nil {
|
||||
if c.closed() {
|
||||
return
|
||||
}
|
||||
var netErr net.Error
|
||||
if goerrors.As(err, &netErr) && netErr.Timeout() {
|
||||
select {
|
||||
@@ -220,9 +217,10 @@ func (c *xicmpConnServer) recv6() {
|
||||
case <-c.closeCh:
|
||||
return
|
||||
}
|
||||
continue
|
||||
}
|
||||
errors.LogErrorInner(context.Background(), err, "recv6 err")
|
||||
continue
|
||||
errors.LogErrorInner(context.Background(), err, "recv err 6")
|
||||
return
|
||||
}
|
||||
|
||||
msg, err := icmp.ParseMessage(58, b[:n])
|
||||
@@ -341,7 +339,6 @@ func (c *xicmpConnServer) Close() error {
|
||||
close(c.closeCh)
|
||||
_ = c.icmp4.Close()
|
||||
_ = c.icmp6.Close()
|
||||
_ = c.conn.Close()
|
||||
c.wg.Wait()
|
||||
select {
|
||||
case p := <-c.readCh:
|
||||
@@ -355,7 +352,7 @@ func (c *xicmpConnServer) Close() error {
|
||||
}
|
||||
|
||||
func (c *xicmpConnServer) LocalAddr() net.Addr {
|
||||
return c.conn.LocalAddr()
|
||||
return &net.UDPAddr{IP: []byte{0, 0, 0, 0}}
|
||||
}
|
||||
|
||||
func (c *xicmpConnServer) SetDeadline(t time.Time) error {
|
||||
|
||||
@@ -2,10 +2,12 @@ package xmc
|
||||
|
||||
import (
|
||||
"fmt"
|
||||
"net"
|
||||
|
||||
"github.com/xtls/xray-core/common/net"
|
||||
"github.com/xtls/xray-core/transport/internet/finalmask"
|
||||
)
|
||||
|
||||
func (c *Config) WrapConnClient(conn net.Conn) (net.Conn, error) {
|
||||
func (c *Config) WrapConnClient(conn net.Conn, dest *net.Destination, dialer *finalmask.Dialer) (net.Conn, error) {
|
||||
profiles, err := profilesFromConfig(c.Profiles)
|
||||
if err != nil {
|
||||
return nil, fmt.Errorf("minecraft finalmask: %w", err)
|
||||
|
||||
@@ -83,7 +83,6 @@ func getGrpcClient(ctx context.Context, dest net.Destination, streamSettings *in
|
||||
}
|
||||
tlsConfig := tls.ConfigFromStreamSettings(streamSettings)
|
||||
realityConfig := reality.ConfigFromStreamSettings(streamSettings)
|
||||
sockopt := streamSettings.SocketSettings
|
||||
grpcSettings := streamSettings.ProtocolSettings.(*Config)
|
||||
|
||||
if client, found := globalDialerMap[dialerConf{dest, streamSettings}]; found && client.GetState() != connectivity.Shutdown {
|
||||
@@ -124,17 +123,13 @@ func getGrpcClient(ctx context.Context, dest net.Destination, streamSettings *in
|
||||
gctx = session.ContextWithOutbounds(gctx, session.OutboundsFromContext(ctx))
|
||||
gctx = session.ContextWithTimeoutOnly(gctx, true)
|
||||
|
||||
c, err := internet.DialSystem(gctx, net.TCPDestination(address, port), sockopt)
|
||||
var c net.Conn
|
||||
if streamSettings.FinalMask != nil {
|
||||
c, err = streamSettings.FinalMask.DialTCP(gctx, net.TCPDestination(address, port))
|
||||
} else {
|
||||
c, err = internet.DialSystem(ctx, dest, streamSettings.SocketSettings)
|
||||
}
|
||||
if err == nil {
|
||||
if streamSettings.TcpmaskManager != nil {
|
||||
newConn, err := streamSettings.TcpmaskManager.WrapConnClient(c)
|
||||
if err != nil {
|
||||
c.Close()
|
||||
return nil, errors.New("mask err").Base(err)
|
||||
}
|
||||
c = newConn
|
||||
}
|
||||
|
||||
if tlsConfig != nil {
|
||||
config := tlsConfig.GetTLSConfig(tls.WithDestination(dest))
|
||||
if fingerprint := tls.GetFingerprint(tlsConfig.Fingerprint); fingerprint != nil {
|
||||
|
||||
@@ -104,28 +104,20 @@ func Listen(ctx context.Context, address net.Address, port net.Port, settings *i
|
||||
go func() {
|
||||
var streamListener net.Listener
|
||||
var err error
|
||||
var addr net.Addr
|
||||
if port == net.Port(0) { // unix
|
||||
streamListener, err = internet.ListenSystem(ctx, &net.UnixAddr{
|
||||
Name: address.Domain(),
|
||||
Net: "unix",
|
||||
}, settings.SocketSettings)
|
||||
if err != nil {
|
||||
errors.LogErrorInner(ctx, err, "failed to listen on ", address)
|
||||
return
|
||||
}
|
||||
addr = &net.UnixAddr{Name: address.Domain(), Net: "unix"}
|
||||
} else { // tcp
|
||||
streamListener, err = internet.ListenSystem(ctx, &net.TCPAddr{
|
||||
IP: address.IP(),
|
||||
Port: int(port),
|
||||
}, settings.SocketSettings)
|
||||
if err != nil {
|
||||
errors.LogErrorInner(ctx, err, "failed to listen on ", address, ":", port)
|
||||
return
|
||||
}
|
||||
addr = &net.TCPAddr{IP: address.IP(), Port: int(port)}
|
||||
}
|
||||
|
||||
if settings.TcpmaskManager != nil {
|
||||
streamListener, _ = settings.TcpmaskManager.WrapListener(streamListener)
|
||||
if settings.FinalMask != nil {
|
||||
streamListener, err = settings.FinalMask.Listen(ctx, addr)
|
||||
} else {
|
||||
streamListener, err = internet.ListenSystem(ctx, addr, settings.SocketSettings)
|
||||
}
|
||||
if err != nil {
|
||||
errors.LogErrorInner(ctx, err, "failed to listen on ", address, ":", port)
|
||||
return
|
||||
}
|
||||
|
||||
errors.LogDebug(ctx, "gRPC listen for service name `"+grpcSettings.getServiceName()+"` tun `"+grpcSettings.getTunStreamName()+"` multi tun `"+grpcSettings.getTunMultiStreamName()+"`")
|
||||
|
||||
@@ -46,21 +46,18 @@ func (c *ConnRF) Read(b []byte) (int, error) {
|
||||
func dialhttpUpgrade(ctx context.Context, dest net.Destination, streamSettings *internet.MemoryStreamConfig) (net.Conn, error) {
|
||||
transportConfiguration := streamSettings.ProtocolSettings.(*Config)
|
||||
|
||||
pconn, err := internet.DialSystem(ctx, dest, streamSettings.SocketSettings)
|
||||
var pconn net.Conn
|
||||
var err error
|
||||
if streamSettings.FinalMask != nil {
|
||||
pconn, err = streamSettings.FinalMask.DialTCP(ctx, dest)
|
||||
} else {
|
||||
pconn, err = internet.DialSystem(ctx, dest, streamSettings.SocketSettings)
|
||||
}
|
||||
if err != nil {
|
||||
errors.LogErrorInner(ctx, err, "failed to dial to ", dest)
|
||||
return nil, err
|
||||
}
|
||||
|
||||
if streamSettings.TcpmaskManager != nil {
|
||||
newConn, err := streamSettings.TcpmaskManager.WrapConnClient(pconn)
|
||||
if err != nil {
|
||||
pconn.Close()
|
||||
return nil, errors.New("mask err").Base(err)
|
||||
}
|
||||
pconn = newConn
|
||||
}
|
||||
|
||||
var conn net.Conn
|
||||
var requestURL url.URL
|
||||
tConfig := tls.ConfigFromStreamSettings(streamSettings)
|
||||
|
||||
@@ -124,29 +124,21 @@ func ListenHTTPUpgrade(ctx context.Context, address net.Address, port net.Port,
|
||||
}
|
||||
var listener net.Listener
|
||||
var err error
|
||||
var addr net.Addr
|
||||
if port == net.Port(0) { // unix
|
||||
listener, err = internet.ListenSystem(ctx, &net.UnixAddr{
|
||||
Name: address.Domain(),
|
||||
Net: "unix",
|
||||
}, streamSettings.SocketSettings)
|
||||
if err != nil {
|
||||
return nil, errors.New("failed to listen unix domain socket(for HttpUpgrade) on ", address).Base(err)
|
||||
}
|
||||
errors.LogInfo(ctx, "listening unix domain socket(for HttpUpgrade) on ", address)
|
||||
addr = &net.UnixAddr{Name: address.Domain(), Net: "unix"}
|
||||
} else { // tcp
|
||||
listener, err = internet.ListenSystem(ctx, &net.TCPAddr{
|
||||
IP: address.IP(),
|
||||
Port: int(port),
|
||||
}, streamSettings.SocketSettings)
|
||||
if err != nil {
|
||||
return nil, errors.New("failed to listen TCP(for HttpUpgrade) on ", address, ":", port).Base(err)
|
||||
}
|
||||
errors.LogInfo(ctx, "listening TCP(for HttpUpgrade) on ", address, ":", port)
|
||||
addr = &net.TCPAddr{IP: address.IP(), Port: int(port)}
|
||||
}
|
||||
|
||||
if streamSettings.TcpmaskManager != nil {
|
||||
listener, _ = streamSettings.TcpmaskManager.WrapListener(listener)
|
||||
if streamSettings.FinalMask != nil {
|
||||
listener, err = streamSettings.FinalMask.Listen(ctx, addr)
|
||||
} else {
|
||||
listener, err = internet.ListenSystem(ctx, addr, streamSettings.SocketSettings)
|
||||
}
|
||||
if err != nil {
|
||||
return nil, errors.New("failed to listen ", addr.Network(), "(for HttpUpgrade) on ", address, ":", port).Base(err)
|
||||
}
|
||||
errors.LogInfo(ctx, "listening ", addr.Network(), "(for HttpUpgrade) on ", address, ":", port)
|
||||
|
||||
if streamSettings.SocketSettings != nil && streamSettings.SocketSettings.AcceptProxyProtocol {
|
||||
errors.LogWarning(ctx, "accepting PROXY protocol")
|
||||
|
||||
@@ -2,7 +2,7 @@ package hysteria
|
||||
|
||||
import (
|
||||
"context"
|
||||
go_tls "crypto/tls"
|
||||
gotls "crypto/tls"
|
||||
"net/http"
|
||||
"net/url"
|
||||
"reflect"
|
||||
@@ -28,12 +28,12 @@ import (
|
||||
type client struct {
|
||||
sync.Mutex
|
||||
|
||||
dest net.Destination
|
||||
config *Config
|
||||
tlsConfig *go_tls.Config
|
||||
socketConfig *internet.SocketConfig
|
||||
udpmaskManager *finalmask.UdpmaskManager
|
||||
quicParams *internet.QuicParams
|
||||
dest net.Destination
|
||||
config *Config
|
||||
tlsConfig *gotls.Config
|
||||
socketConfig *internet.SocketConfig
|
||||
finalMask *finalmask.FinalMask
|
||||
quicParams *internet.QuicParams
|
||||
|
||||
conn *quic.Conn
|
||||
tr *quic.Transport
|
||||
@@ -113,30 +113,29 @@ func (c *client) dial(ctx context.Context) error {
|
||||
// }
|
||||
|
||||
var pktConn net.PacketConn
|
||||
var udpAddr *net.UDPAddr
|
||||
|
||||
raw, err := internet.DialSystem(ctx, c.dest, c.socketConfig)
|
||||
if err != nil {
|
||||
return errors.New("failed to dial to dest").Base(err)
|
||||
}
|
||||
switch c := raw.(type) {
|
||||
case *internet.PacketConnWrapper:
|
||||
pktConn = c.PacketConn
|
||||
udpAddr = raw.RemoteAddr().(*net.UDPAddr)
|
||||
case *cnc.Connection:
|
||||
pktConn = &internet.FakePacketConn{Conn: c}
|
||||
udpAddr = &net.UDPAddr{IP: c.RemoteAddr().(*net.TCPAddr).IP, Port: c.RemoteAddr().(*net.TCPAddr).Port}
|
||||
default:
|
||||
panic(reflect.TypeOf(c))
|
||||
}
|
||||
|
||||
if c.udpmaskManager != nil {
|
||||
newConn, err := c.udpmaskManager.WrapPacketConnClient(pktConn)
|
||||
var udpAddr net.Addr
|
||||
if c.finalMask != nil {
|
||||
conn, err := c.finalMask.DialUDP(ctx, c.dest)
|
||||
if err != nil {
|
||||
pktConn.Close()
|
||||
return errors.New("mask err").Base(err)
|
||||
return errors.New("failed to dial to dest").Base(err)
|
||||
}
|
||||
pktConn = conn.(*finalmask.PacketConnWrapper).PacketConn
|
||||
udpAddr = conn.RemoteAddr()
|
||||
} else {
|
||||
conn, err := internet.DialSystem(ctx, c.dest, c.socketConfig)
|
||||
if err != nil {
|
||||
return errors.New("failed to dial to dest").Base(err)
|
||||
}
|
||||
switch c := conn.(type) {
|
||||
case *internet.PacketConnWrapper:
|
||||
pktConn = c.PacketConn
|
||||
udpAddr = c.RemoteAddr()
|
||||
case *cnc.Connection:
|
||||
pktConn = &internet.FakePacketConn{Conn: c}
|
||||
udpAddr = &net.UDPAddr{IP: []byte{0, 0, 0, 0}}
|
||||
default:
|
||||
panic(reflect.TypeOf(c))
|
||||
}
|
||||
pktConn = newConn
|
||||
}
|
||||
|
||||
tr := &quic.Transport{Conn: pktConn, DisableGSO: quicParams.DisableGSO}
|
||||
@@ -150,7 +149,7 @@ func (c *client) dial(ctx context.Context) error {
|
||||
rt := &http3.Transport{
|
||||
TLSClientConfig: c.tlsConfig,
|
||||
QUICConfig: quicConfig,
|
||||
Dial: func(ctx context.Context, _ string, tlsCfg *go_tls.Config, cfg *quic.Config) (*quic.Conn, error) {
|
||||
Dial: func(ctx context.Context, _ string, tlsCfg *gotls.Config, cfg *quic.Config) (*quic.Conn, error) {
|
||||
qc, err := tr.DialEarly(ctx, udpAddr, tlsCfg, cfg)
|
||||
if err != nil {
|
||||
return nil, err
|
||||
@@ -316,12 +315,12 @@ func Dial(ctx context.Context, dest net.Destination, streamSettings *internet.Me
|
||||
c = manager.m[dialerConf{dest, streamSettings}]
|
||||
if c == nil {
|
||||
c = &client{
|
||||
dest: dest,
|
||||
config: streamSettings.ProtocolSettings.(*Config),
|
||||
tlsConfig: tlsConfig.GetTLSConfig(tls.WithDestination(dest)),
|
||||
socketConfig: streamSettings.SocketSettings,
|
||||
udpmaskManager: streamSettings.UdpmaskManager,
|
||||
quicParams: streamSettings.QuicParams,
|
||||
dest: dest,
|
||||
config: streamSettings.ProtocolSettings.(*Config),
|
||||
tlsConfig: tlsConfig.GetTLSConfig(tls.WithDestination(dest)),
|
||||
socketConfig: streamSettings.SocketSettings,
|
||||
finalMask: streamSettings.FinalMask,
|
||||
quicParams: streamSettings.QuicParams,
|
||||
}
|
||||
manager.m[dialerConf{dest, streamSettings}] = c
|
||||
}
|
||||
|
||||
@@ -316,20 +316,17 @@ func Listen(ctx context.Context, address net.Address, port net.Port, streamSetti
|
||||
quicConfig.MaxIncomingStreams = 1024
|
||||
}
|
||||
|
||||
pktConn, err := internet.ListenSystemPacket(context.Background(), &net.UDPAddr{IP: address.IP(), Port: int(port)}, streamSettings.SocketSettings)
|
||||
var pktConn net.PacketConn
|
||||
var err error
|
||||
if streamSettings.FinalMask != nil {
|
||||
pktConn, err = streamSettings.FinalMask.ListenPacket(context.Background(), &net.UDPAddr{IP: address.IP(), Port: int(port)})
|
||||
} else {
|
||||
pktConn, err = internet.ListenSystemPacket(context.Background(), &net.UDPAddr{IP: address.IP(), Port: int(port)}, streamSettings.SocketSettings)
|
||||
}
|
||||
if err != nil {
|
||||
return nil, err
|
||||
}
|
||||
|
||||
if streamSettings.UdpmaskManager != nil {
|
||||
newConn, err := streamSettings.UdpmaskManager.WrapPacketConnServer(pktConn)
|
||||
if err != nil {
|
||||
pktConn.Close()
|
||||
return nil, errors.New("mask err").Base(err)
|
||||
}
|
||||
pktConn = newConn
|
||||
}
|
||||
|
||||
var k *quic.StatelessResetKey
|
||||
if !quicParams.DisableStatelessReset {
|
||||
k = &quic.StatelessResetKey{}
|
||||
|
||||
@@ -0,0 +1,165 @@
|
||||
package hysteria
|
||||
|
||||
import (
|
||||
"context"
|
||||
"crypto/tls"
|
||||
"crypto/x509"
|
||||
"errors"
|
||||
"net"
|
||||
"runtime"
|
||||
"testing"
|
||||
"time"
|
||||
|
||||
"github.com/apernet/quic-go"
|
||||
"github.com/xtls/xray-core/common"
|
||||
"github.com/xtls/xray-core/common/protocol/tls/cert"
|
||||
)
|
||||
|
||||
func TestDatagram(t *testing.T) {
|
||||
run := func() (addr net.Addr, recv chan int64, cancel func()) {
|
||||
cert, _ := cert.MustGenerate(nil)
|
||||
Certificate := [][]byte{cert.Certificate}
|
||||
PrivateKey := common.Must2(x509.ParsePKCS8PrivateKey(cert.PrivateKey))
|
||||
|
||||
tlsConf := &tls.Config{
|
||||
Certificates: []tls.Certificate{
|
||||
{
|
||||
Certificate: Certificate,
|
||||
PrivateKey: PrivateKey,
|
||||
},
|
||||
},
|
||||
NextProtos: []string{"h3"},
|
||||
}
|
||||
|
||||
quicConf := &quic.Config{
|
||||
InitialStreamReceiveWindow: 8388608,
|
||||
MaxStreamReceiveWindow: 8388608,
|
||||
InitialConnectionReceiveWindow: 8388608 * 5 / 2,
|
||||
MaxConnectionReceiveWindow: 8388608 * 5 / 2,
|
||||
MaxIdleTimeout: 30 * time.Second,
|
||||
MaxIncomingStreams: 1024,
|
||||
DisablePathMTUDiscovery: runtime.GOOS != "linux" && runtime.GOOS != "windows" && runtime.GOOS != "darwin",
|
||||
EnableDatagrams: true,
|
||||
MaxDatagramFrameSize: MaxDatagramFrameSize,
|
||||
AssumePeerMaxDatagramFrameSize: MaxDatagramFrameSize,
|
||||
DisablePathManager: true,
|
||||
}
|
||||
|
||||
pktConn := common.Must2(net.ListenPacket("udp", "127.0.0.1:0"))
|
||||
tr := &quic.Transport{Conn: pktConn}
|
||||
l := common.Must2(tr.Listen(tlsConf, quicConf))
|
||||
|
||||
recv = make(chan int64)
|
||||
ctx, cancel := context.WithCancel(context.Background())
|
||||
|
||||
go func() {
|
||||
defer pktConn.Close()
|
||||
defer tr.Close()
|
||||
defer l.Close()
|
||||
defer close(recv)
|
||||
|
||||
var buf [1500]byte
|
||||
for {
|
||||
conn, err := l.Accept(ctx)
|
||||
if err != nil {
|
||||
if !errors.Is(err, context.Canceled) {
|
||||
t.Error(err)
|
||||
}
|
||||
break
|
||||
}
|
||||
err = conn.SendDatagram(buf[:])
|
||||
var qErr *quic.DatagramTooLargeError
|
||||
if !errors.As(err, &qErr) {
|
||||
t.Error(err)
|
||||
}
|
||||
recv <- qErr.MaxDatagramPayloadSize
|
||||
defer conn.CloseWithError(0, "")
|
||||
}
|
||||
}()
|
||||
|
||||
return l.Addr(), recv, cancel
|
||||
}
|
||||
|
||||
addr, recv, cancel := run()
|
||||
|
||||
t.Run("With ChromeParrot", func(t *testing.T) {
|
||||
tlsConf := &tls.Config{
|
||||
InsecureSkipVerify: true,
|
||||
}
|
||||
|
||||
quicConf := &quic.Config{
|
||||
InitialStreamReceiveWindow: 8388608,
|
||||
MaxStreamReceiveWindow: 8388608,
|
||||
InitialConnectionReceiveWindow: 8388608 * 5 / 2,
|
||||
MaxConnectionReceiveWindow: 8388608 * 5 / 2,
|
||||
MaxIdleTimeout: 30 * time.Second,
|
||||
KeepAlivePeriod: 10 * time.Second,
|
||||
DisablePathMTUDiscovery: runtime.GOOS != "linux" && runtime.GOOS != "windows" && runtime.GOOS != "darwin",
|
||||
ChromeParrot: true,
|
||||
EnableDatagrams: true,
|
||||
MaxDatagramFrameSize: MaxDatagramFrameSize,
|
||||
OmitMaxDatagramFrameSize: true,
|
||||
DisablePathManager: true,
|
||||
}
|
||||
|
||||
pktConn := common.Must2(net.ListenPacket("udp", "127.0.0.1:0"))
|
||||
tr := &quic.Transport{Conn: pktConn, ConnectionIDGenerator: quic.ZeroLengthConnectionIDGenerator{}}
|
||||
conn := common.Must2(tr.DialEarly(context.Background(), addr, tlsConf, quicConf))
|
||||
|
||||
defer pktConn.Close()
|
||||
defer tr.Close()
|
||||
defer conn.CloseWithError(0, "")
|
||||
|
||||
var buf [1500]byte
|
||||
err := conn.SendDatagram(buf[:])
|
||||
var qErr *quic.DatagramTooLargeError
|
||||
if !errors.As(err, &qErr) || qErr.MaxDatagramPayloadSize != 1197 {
|
||||
t.Error(err)
|
||||
}
|
||||
if server := <-recv; server != 1243 {
|
||||
t.Error(server)
|
||||
}
|
||||
})
|
||||
|
||||
t.Run("Without ChromeParrot", func(t *testing.T) {
|
||||
tlsConf := &tls.Config{
|
||||
InsecureSkipVerify: true,
|
||||
NextProtos: []string{"h3"},
|
||||
}
|
||||
|
||||
quicConf := &quic.Config{
|
||||
InitialStreamReceiveWindow: 8388608,
|
||||
MaxStreamReceiveWindow: 8388608,
|
||||
InitialConnectionReceiveWindow: 8388608 * 5 / 2,
|
||||
MaxConnectionReceiveWindow: 8388608 * 5 / 2,
|
||||
MaxIdleTimeout: 30 * time.Second,
|
||||
KeepAlivePeriod: 10 * time.Second,
|
||||
DisablePathMTUDiscovery: runtime.GOOS != "linux" && runtime.GOOS != "windows" && runtime.GOOS != "darwin",
|
||||
ChromeParrot: false,
|
||||
EnableDatagrams: true,
|
||||
MaxDatagramFrameSize: MaxDatagramFrameSize,
|
||||
OmitMaxDatagramFrameSize: true,
|
||||
DisablePathManager: true,
|
||||
}
|
||||
|
||||
pktConn := common.Must2(net.ListenPacket("udp", "127.0.0.1:0"))
|
||||
tr := &quic.Transport{Conn: pktConn}
|
||||
conn := common.Must2(tr.DialEarly(context.Background(), addr, tlsConf, quicConf))
|
||||
|
||||
defer pktConn.Close()
|
||||
defer tr.Close()
|
||||
defer conn.CloseWithError(0, "")
|
||||
|
||||
var buf [1500]byte
|
||||
err := conn.SendDatagram(buf[:])
|
||||
var qErr *quic.DatagramTooLargeError
|
||||
if !errors.As(err, &qErr) || qErr.MaxDatagramPayloadSize != 1197 {
|
||||
t.Error(err)
|
||||
}
|
||||
if server := <-recv; server != 1197 {
|
||||
t.Error(server)
|
||||
}
|
||||
})
|
||||
|
||||
cancel()
|
||||
}
|
||||
@@ -3,7 +3,6 @@ package kcp
|
||||
import (
|
||||
"context"
|
||||
"io"
|
||||
reflect "reflect"
|
||||
"sync/atomic"
|
||||
|
||||
"github.com/xtls/xray-core/common"
|
||||
@@ -11,7 +10,6 @@ import (
|
||||
"github.com/xtls/xray-core/common/dice"
|
||||
"github.com/xtls/xray-core/common/errors"
|
||||
"github.com/xtls/xray-core/common/net"
|
||||
"github.com/xtls/xray-core/common/net/cnc"
|
||||
"github.com/xtls/xray-core/transport/internet"
|
||||
"github.com/xtls/xray-core/transport/internet/stat"
|
||||
"github.com/xtls/xray-core/transport/internet/tls"
|
||||
@@ -51,36 +49,17 @@ func DialKCP(ctx context.Context, dest net.Destination, streamSettings *internet
|
||||
dest.Network = net.Network_UDP
|
||||
errors.LogInfo(ctx, "dialing mKCP to ", dest)
|
||||
|
||||
conn, err := internet.DialSystem(ctx, dest, streamSettings.SocketSettings)
|
||||
var conn net.Conn
|
||||
var err error
|
||||
if streamSettings.FinalMask != nil {
|
||||
conn, err = streamSettings.FinalMask.DialUDP(ctx, dest)
|
||||
} else {
|
||||
conn, err = internet.DialSystem(ctx, dest, streamSettings.SocketSettings)
|
||||
}
|
||||
if err != nil {
|
||||
return nil, errors.New("failed to dial to dest: ", err).AtWarning().Base(err)
|
||||
}
|
||||
|
||||
if streamSettings.UdpmaskManager != nil {
|
||||
var pktConn net.PacketConn
|
||||
var udpAddr *net.UDPAddr
|
||||
switch c := conn.(type) {
|
||||
case *internet.PacketConnWrapper:
|
||||
pktConn = c.PacketConn
|
||||
udpAddr = c.RemoteAddr().(*net.UDPAddr)
|
||||
case *cnc.Connection:
|
||||
pktConn = &internet.FakePacketConn{Conn: c}
|
||||
udpAddr = &net.UDPAddr{IP: c.RemoteAddr().(*net.TCPAddr).IP, Port: c.RemoteAddr().(*net.TCPAddr).Port}
|
||||
default:
|
||||
panic(reflect.TypeOf(c))
|
||||
}
|
||||
newConn, err := streamSettings.UdpmaskManager.WrapPacketConnClient(pktConn)
|
||||
if err != nil {
|
||||
pktConn.Close()
|
||||
return nil, errors.New("mask err").Base(err)
|
||||
}
|
||||
pktConn = newConn
|
||||
conn = &internet.PacketConnWrapper{
|
||||
PacketConn: pktConn,
|
||||
Dest: udpAddr,
|
||||
}
|
||||
}
|
||||
|
||||
kcpSettings := streamSettings.ProtocolSettings.(*Config)
|
||||
|
||||
reader := &KCPPacketReader{}
|
||||
|
||||
@@ -1,7 +1,12 @@
|
||||
package internet
|
||||
|
||||
import (
|
||||
"context"
|
||||
"reflect"
|
||||
|
||||
"github.com/xtls/xray-core/common"
|
||||
"github.com/xtls/xray-core/common/net"
|
||||
"github.com/xtls/xray-core/common/net/cnc"
|
||||
"github.com/xtls/xray-core/transport/internet/finalmask"
|
||||
)
|
||||
|
||||
@@ -12,8 +17,7 @@ type MemoryStreamConfig struct {
|
||||
ProtocolSettings interface{}
|
||||
SecurityType string
|
||||
SecuritySettings interface{}
|
||||
TcpmaskManager *finalmask.TcpmaskManager
|
||||
UdpmaskManager *finalmask.UdpmaskManager
|
||||
FinalMask *finalmask.FinalMask
|
||||
QuicParams *QuicParams
|
||||
SocketSettings *SocketConfig
|
||||
DownloadSettings *MemoryStreamConfig
|
||||
@@ -51,33 +55,53 @@ func ToMemoryStreamConfig(s *StreamConfig) (*MemoryStreamConfig, error) {
|
||||
mss.SecuritySettings = ess
|
||||
}
|
||||
|
||||
if s != nil && len(s.Tcpmasks) > 0 {
|
||||
var masks []finalmask.Tcpmask
|
||||
for _, msg := range s.Tcpmasks {
|
||||
instance, err := msg.GetInstance()
|
||||
if err != nil {
|
||||
return nil, err
|
||||
}
|
||||
masks = append(masks, instance.(finalmask.Tcpmask))
|
||||
var tcpMasks []finalmask.TCPMask
|
||||
var udpMasks []finalmask.UDPMask
|
||||
|
||||
if s != nil {
|
||||
for i := range s.Tcpmasks {
|
||||
instance := common.Must2(s.Tcpmasks[i].GetInstance())
|
||||
tcpMasks = append(tcpMasks, instance.(finalmask.TCPMask))
|
||||
}
|
||||
for i := range s.Udpmasks {
|
||||
instance := common.Must2(s.Udpmasks[i].GetInstance())
|
||||
udpMasks = append(udpMasks, instance.(finalmask.UDPMask))
|
||||
}
|
||||
mss.TcpmaskManager = finalmask.NewTcpmaskManager(masks)
|
||||
}
|
||||
|
||||
dialTCP := func(ctx context.Context, dest net.Destination) (net.Conn, error) {
|
||||
return DialSystem(ctx, dest, mss.SocketSettings)
|
||||
}
|
||||
listen := func(ctx context.Context, addr net.Addr) (net.Listener, error) {
|
||||
return ListenSystem(ctx, addr, mss.SocketSettings)
|
||||
}
|
||||
dialUDP := func(ctx context.Context, dest net.Destination) (net.PacketConn, net.Addr, error) {
|
||||
conn, err := DialSystem(ctx, dest, mss.SocketSettings)
|
||||
if err != nil {
|
||||
return nil, nil, err
|
||||
}
|
||||
var newConn net.PacketConn
|
||||
var udpAddr net.Addr
|
||||
switch c := conn.(type) {
|
||||
case *PacketConnWrapper:
|
||||
newConn = c.PacketConn
|
||||
udpAddr = conn.RemoteAddr()
|
||||
case *cnc.Connection:
|
||||
newConn = &FakePacketConn{Conn: c}
|
||||
udpAddr = &net.UDPAddr{IP: []byte{0, 0, 0, 0}, Port: 0}
|
||||
default:
|
||||
panic(reflect.TypeOf(c))
|
||||
}
|
||||
return newConn, udpAddr, nil
|
||||
}
|
||||
listenPacket := func(ctx context.Context, addr net.Addr) (net.PacketConn, error) {
|
||||
return ListenSystemPacket(ctx, addr, mss.SocketSettings)
|
||||
}
|
||||
mss.FinalMask = finalmask.NewFinalMask(tcpMasks, udpMasks, dialTCP, listen, dialUDP, listenPacket)
|
||||
|
||||
if s != nil && s.QuicParams != nil {
|
||||
mss.QuicParams = s.QuicParams
|
||||
}
|
||||
|
||||
if s != nil && len(s.Udpmasks) > 0 {
|
||||
var masks []finalmask.Udpmask
|
||||
for _, msg := range s.Udpmasks {
|
||||
instance, err := msg.GetInstance()
|
||||
if err != nil {
|
||||
return nil, err
|
||||
}
|
||||
masks = append(masks, instance.(finalmask.Udpmask))
|
||||
}
|
||||
mss.UdpmaskManager = finalmask.NewUdpmaskManager(masks)
|
||||
}
|
||||
|
||||
return mss, nil
|
||||
}
|
||||
|
||||
@@ -8,7 +8,7 @@ import (
|
||||
"net/http"
|
||||
"net/http/httptrace"
|
||||
"net/url"
|
||||
reflect "reflect"
|
||||
"reflect"
|
||||
"runtime"
|
||||
"strconv"
|
||||
"sync"
|
||||
@@ -25,6 +25,7 @@ import (
|
||||
"github.com/xtls/xray-core/common/signal/done"
|
||||
"github.com/xtls/xray-core/transport/internet"
|
||||
"github.com/xtls/xray-core/transport/internet/browser_dialer"
|
||||
"github.com/xtls/xray-core/transport/internet/finalmask"
|
||||
"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/reality"
|
||||
@@ -116,20 +117,17 @@ func createHTTPClient(dest net.Destination, streamSettings *internet.MemoryStrea
|
||||
transportConfig := streamSettings.ProtocolSettings.(*Config)
|
||||
|
||||
dialContext := func(ctxInner context.Context) (net.Conn, error) {
|
||||
conn, err := internet.DialSystem(ctxInner, dest, streamSettings.SocketSettings)
|
||||
var conn net.Conn
|
||||
var err error
|
||||
if streamSettings.FinalMask != nil {
|
||||
conn, err = streamSettings.FinalMask.DialTCP(ctxInner, dest)
|
||||
} else {
|
||||
conn, err = internet.DialSystem(ctxInner, dest, streamSettings.SocketSettings)
|
||||
}
|
||||
if err != nil {
|
||||
return nil, err
|
||||
}
|
||||
|
||||
if streamSettings.TcpmaskManager != nil {
|
||||
newConn, err := streamSettings.TcpmaskManager.WrapConnClient(conn)
|
||||
if err != nil {
|
||||
conn.Close()
|
||||
return nil, errors.New("mask err").Base(err)
|
||||
}
|
||||
conn = newConn
|
||||
}
|
||||
|
||||
if realityConfig != nil {
|
||||
return reality.UClient(conn, realityConfig, ctxInner, dest)
|
||||
}
|
||||
@@ -196,30 +194,29 @@ func createHTTPClient(dest net.Destination, streamSettings *internet.MemoryStrea
|
||||
TLSClientConfig: gotlsConfig,
|
||||
Dial: func(ctx context.Context, addr string, tlsCfg *gotls.Config, cfg *quic.Config) (*quic.Conn, error) {
|
||||
var pktConn net.PacketConn
|
||||
var udpAddr *net.UDPAddr
|
||||
|
||||
raw, err := internet.DialSystem(ctx, dest, streamSettings.SocketSettings)
|
||||
if err != nil {
|
||||
return nil, errors.New("failed to dial to dest").Base(err)
|
||||
}
|
||||
switch c := raw.(type) {
|
||||
case *internet.PacketConnWrapper:
|
||||
pktConn = c.PacketConn
|
||||
udpAddr = raw.RemoteAddr().(*net.UDPAddr)
|
||||
case *cnc.Connection:
|
||||
pktConn = &internet.FakePacketConn{Conn: c}
|
||||
udpAddr = &net.UDPAddr{IP: c.RemoteAddr().(*net.TCPAddr).IP, Port: c.RemoteAddr().(*net.TCPAddr).Port}
|
||||
default:
|
||||
panic(reflect.TypeOf(c))
|
||||
}
|
||||
|
||||
if streamSettings.UdpmaskManager != nil {
|
||||
newConn, err := streamSettings.UdpmaskManager.WrapPacketConnClient(pktConn)
|
||||
var udpAddr net.Addr
|
||||
if streamSettings.FinalMask != nil {
|
||||
conn, err := streamSettings.FinalMask.DialUDP(ctx, dest)
|
||||
if err != nil {
|
||||
pktConn.Close()
|
||||
return nil, errors.New("mask err").Base(err)
|
||||
return nil, errors.New("failed to dial to dest").Base(err)
|
||||
}
|
||||
pktConn = conn.(*finalmask.PacketConnWrapper).PacketConn
|
||||
udpAddr = conn.RemoteAddr()
|
||||
} else {
|
||||
conn, err := internet.DialSystem(ctx, dest, streamSettings.SocketSettings)
|
||||
if err != nil {
|
||||
return nil, errors.New("failed to dial to dest").Base(err)
|
||||
}
|
||||
switch c := conn.(type) {
|
||||
case *internet.PacketConnWrapper:
|
||||
pktConn = c.PacketConn
|
||||
udpAddr = c.RemoteAddr()
|
||||
case *cnc.Connection:
|
||||
pktConn = &internet.FakePacketConn{Conn: c}
|
||||
udpAddr = &net.UDPAddr{IP: []byte{0, 0, 0, 0}}
|
||||
default:
|
||||
panic(reflect.TypeOf(c))
|
||||
}
|
||||
pktConn = newConn
|
||||
}
|
||||
|
||||
tr := &quic.Transport{Conn: pktConn, DisableGSO: quicParams.DisableGSO}
|
||||
|
||||
@@ -463,31 +463,17 @@ func ListenXH(ctx context.Context, address net.Address, port net.Port, streamSet
|
||||
l.isH3 = len(tlsConfig.NextProtos) == 1 && tlsConfig.NextProtos[0] == "h3"
|
||||
|
||||
var err error
|
||||
if port == net.Port(0) { // unix
|
||||
l.listener, err = internet.ListenSystem(ctx, &net.UnixAddr{
|
||||
Name: address.Domain(),
|
||||
Net: "unix",
|
||||
}, streamSettings.SocketSettings)
|
||||
if err != nil {
|
||||
return nil, errors.New("failed to listen UNIX domain socket for XHTTP on ", address).Base(err)
|
||||
if l.isH3 {
|
||||
var pktConn net.PacketConn
|
||||
var err error
|
||||
if streamSettings.FinalMask != nil {
|
||||
pktConn, err = streamSettings.FinalMask.ListenPacket(context.Background(), &net.UDPAddr{IP: address.IP(), Port: int(port)})
|
||||
} else {
|
||||
pktConn, err = internet.ListenSystemPacket(context.Background(), &net.UDPAddr{IP: address.IP(), Port: int(port)}, streamSettings.SocketSettings)
|
||||
}
|
||||
errors.LogInfo(ctx, "listening UNIX domain socket for XHTTP on ", address)
|
||||
} else if l.isH3 { // quic
|
||||
Conn, err := internet.ListenSystemPacket(context.Background(), &net.UDPAddr{
|
||||
IP: address.IP(),
|
||||
Port: int(port),
|
||||
}, streamSettings.SocketSettings)
|
||||
if err != nil {
|
||||
return nil, errors.New("failed to listen UDP for XHTTP/3 on ", address, ":", port).Base(err)
|
||||
}
|
||||
if streamSettings.UdpmaskManager != nil {
|
||||
newConn, err := streamSettings.UdpmaskManager.WrapPacketConnServer(Conn)
|
||||
if err != nil {
|
||||
Conn.Close()
|
||||
return nil, errors.New("mask err").Base(err)
|
||||
}
|
||||
Conn = newConn
|
||||
}
|
||||
|
||||
quicParams := streamSettings.QuicParams
|
||||
if quicParams == nil {
|
||||
@@ -512,7 +498,7 @@ func ListenXH(ctx context.Context, address net.Address, port net.Port, streamSet
|
||||
common.Must2(rand.Read((*k)[:]))
|
||||
}
|
||||
|
||||
tr := &quic.Transport{Conn: Conn, DisableGSO: quicParams.DisableGSO, StatelessResetKey: k}
|
||||
tr := &quic.Transport{Conn: pktConn, DisableGSO: quicParams.DisableGSO, StatelessResetKey: k}
|
||||
|
||||
l.h3listener, err = tr.ListenEarly(tlsConfig, quicConfig)
|
||||
if err != nil {
|
||||
@@ -534,21 +520,24 @@ func ListenXH(ctx context.Context, address net.Address, port net.Port, streamSet
|
||||
errors.LogErrorInner(ctx, err, "failed to serve HTTP/3 for XHTTP/3")
|
||||
}
|
||||
_ = tr.Close()
|
||||
_ = Conn.Close()
|
||||
_ = pktConn.Close()
|
||||
}()
|
||||
} else { // tcp
|
||||
l.listener, err = internet.ListenSystem(ctx, &net.TCPAddr{
|
||||
IP: address.IP(),
|
||||
Port: int(port),
|
||||
}, streamSettings.SocketSettings)
|
||||
if err != nil {
|
||||
return nil, errors.New("failed to listen TCP for XHTTP on ", address, ":", port).Base(err)
|
||||
} else {
|
||||
var addr net.Addr
|
||||
if port == net.Port(0) { // unix
|
||||
addr = &net.UnixAddr{Name: address.Domain(), Net: "unix"}
|
||||
} else { // tcp
|
||||
addr = &net.TCPAddr{IP: address.IP(), Port: int(port)}
|
||||
}
|
||||
errors.LogInfo(ctx, "listening TCP for XHTTP on ", address, ":", port)
|
||||
}
|
||||
|
||||
if !l.isH3 && streamSettings.TcpmaskManager != nil {
|
||||
l.listener, _ = streamSettings.TcpmaskManager.WrapListener(l.listener)
|
||||
if streamSettings.FinalMask != nil {
|
||||
l.listener, err = streamSettings.FinalMask.Listen(ctx, addr)
|
||||
} else {
|
||||
l.listener, err = internet.ListenSystem(ctx, addr, streamSettings.SocketSettings)
|
||||
}
|
||||
if err != nil {
|
||||
return nil, errors.New("failed to listen ", addr.Network(), " for XHTTP on ", address, ":", port).Base(err)
|
||||
}
|
||||
errors.LogInfo(ctx, "listening ", addr.Network(), " for XHTTP on ", address, ":", port)
|
||||
}
|
||||
|
||||
// tcp/unix (h1/h2)
|
||||
|
||||
@@ -235,5 +235,5 @@ func (c *FakePacketConn) WriteTo(p []byte, _ net.Addr) (n int, err error) {
|
||||
}
|
||||
|
||||
func (c *FakePacketConn) LocalAddr() net.Addr {
|
||||
return &net.UDPAddr{IP: c.Conn.LocalAddr().(*net.TCPAddr).IP, Port: c.Conn.LocalAddr().(*net.TCPAddr).Port}
|
||||
return &net.UDPAddr{IP: []byte{0, 0, 0, 0}}
|
||||
}
|
||||
|
||||
@@ -19,18 +19,15 @@ import (
|
||||
// Dial dials a new TCP connection to the given destination.
|
||||
func Dial(ctx context.Context, dest net.Destination, streamSettings *internet.MemoryStreamConfig) (stat.Connection, error) {
|
||||
errors.LogInfo(ctx, "dialing TCP to ", dest)
|
||||
conn, err := internet.DialSystem(ctx, dest, streamSettings.SocketSettings)
|
||||
if err != nil {
|
||||
return nil, err
|
||||
var conn net.Conn
|
||||
var err error
|
||||
if streamSettings.FinalMask != nil {
|
||||
conn, err = streamSettings.FinalMask.DialTCP(ctx, dest)
|
||||
} else {
|
||||
conn, err = internet.DialSystem(ctx, dest, streamSettings.SocketSettings)
|
||||
}
|
||||
|
||||
if streamSettings.TcpmaskManager != nil {
|
||||
newConn, err := streamSettings.TcpmaskManager.WrapConnClient(conn)
|
||||
if err != nil {
|
||||
conn.Close()
|
||||
return nil, errors.New("mask err").Base(err)
|
||||
}
|
||||
conn = newConn
|
||||
if err != nil {
|
||||
return nil, errors.New("failed to dial to dest").Base(err)
|
||||
}
|
||||
|
||||
if config := tls.ConfigFromStreamSettings(streamSettings); config != nil {
|
||||
|
||||
@@ -41,29 +41,21 @@ func ListenTCP(ctx context.Context, address net.Address, port net.Port, streamSe
|
||||
}
|
||||
var listener net.Listener
|
||||
var err error
|
||||
var addr net.Addr
|
||||
if port == net.Port(0) { // unix
|
||||
listener, err = internet.ListenSystem(ctx, &net.UnixAddr{
|
||||
Name: address.Domain(),
|
||||
Net: "unix",
|
||||
}, streamSettings.SocketSettings)
|
||||
if err != nil {
|
||||
return nil, errors.New("failed to listen Unix Domain Socket on ", address).Base(err)
|
||||
}
|
||||
errors.LogInfo(ctx, "listening Unix Domain Socket on ", address)
|
||||
addr = &net.UnixAddr{Name: address.Domain(), Net: "unix"}
|
||||
} else { // tcp
|
||||
addr = &net.TCPAddr{IP: address.IP(), Port: int(port)}
|
||||
}
|
||||
if streamSettings.FinalMask != nil {
|
||||
listener, err = streamSettings.FinalMask.Listen(ctx, addr)
|
||||
} else {
|
||||
listener, err = internet.ListenSystem(ctx, &net.TCPAddr{
|
||||
IP: address.IP(),
|
||||
Port: int(port),
|
||||
}, streamSettings.SocketSettings)
|
||||
if err != nil {
|
||||
return nil, errors.New("failed to listen TCP on ", address, ":", port).Base(err)
|
||||
}
|
||||
errors.LogInfo(ctx, "listening TCP on ", address, ":", port)
|
||||
listener, err = internet.ListenSystem(ctx, addr, streamSettings.SocketSettings)
|
||||
}
|
||||
|
||||
if streamSettings.TcpmaskManager != nil {
|
||||
listener, _ = streamSettings.TcpmaskManager.WrapListener(listener)
|
||||
if err != nil {
|
||||
return nil, errors.New("failed to listen ", addr.Network(), " on ", address, ":", port).Base(err)
|
||||
}
|
||||
errors.LogInfo(ctx, "listening ", addr.Network(), " on ", address, ":", port)
|
||||
|
||||
if streamSettings.SocketSettings != nil && streamSettings.SocketSettings.AcceptProxyProtocol {
|
||||
errors.LogWarning(ctx, "accepting PROXY protocol")
|
||||
|
||||
@@ -2,12 +2,9 @@ package udp
|
||||
|
||||
import (
|
||||
"context"
|
||||
"reflect"
|
||||
|
||||
"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/common/net/cnc"
|
||||
"github.com/xtls/xray-core/transport/internet"
|
||||
"github.com/xtls/xray-core/transport/internet/stat"
|
||||
)
|
||||
@@ -15,40 +12,14 @@ import (
|
||||
func init() {
|
||||
common.Must(internet.RegisterTransportDialer(protocolName,
|
||||
func(ctx context.Context, dest net.Destination, streamSettings *internet.MemoryStreamConfig) (stat.Connection, error) {
|
||||
var sockopt *internet.SocketConfig
|
||||
if streamSettings != nil {
|
||||
sockopt = streamSettings.SocketSettings
|
||||
}
|
||||
conn, err := internet.DialSystem(ctx, dest, sockopt)
|
||||
if err != nil {
|
||||
return nil, err
|
||||
}
|
||||
|
||||
if streamSettings != nil && streamSettings.UdpmaskManager != nil {
|
||||
var pktConn net.PacketConn
|
||||
var udpAddr *net.UDPAddr
|
||||
switch c := conn.(type) {
|
||||
case *internet.PacketConnWrapper:
|
||||
pktConn = c.PacketConn
|
||||
udpAddr = c.RemoteAddr().(*net.UDPAddr)
|
||||
case *cnc.Connection:
|
||||
pktConn = &internet.FakePacketConn{Conn: c}
|
||||
udpAddr = &net.UDPAddr{IP: c.RemoteAddr().(*net.TCPAddr).IP, Port: c.RemoteAddr().(*net.TCPAddr).Port}
|
||||
default:
|
||||
panic(reflect.TypeOf(c))
|
||||
}
|
||||
newConn, err := streamSettings.UdpmaskManager.WrapPacketConnClient(pktConn)
|
||||
if err != nil {
|
||||
pktConn.Close()
|
||||
return nil, errors.New("mask err").Base(err)
|
||||
}
|
||||
pktConn = newConn
|
||||
conn = &internet.PacketConnWrapper{
|
||||
PacketConn: pktConn,
|
||||
Dest: udpAddr,
|
||||
if streamSettings != nil && streamSettings.FinalMask != nil {
|
||||
return streamSettings.FinalMask.DialUDP(ctx, dest)
|
||||
} else {
|
||||
var sockopt *internet.SocketConfig
|
||||
if streamSettings != nil && streamSettings.SocketSettings != nil {
|
||||
sockopt = streamSettings.SocketSettings
|
||||
}
|
||||
return internet.DialSystem(ctx, dest, sockopt)
|
||||
}
|
||||
|
||||
return conn, nil
|
||||
}))
|
||||
}
|
||||
|
||||
@@ -58,24 +58,15 @@ func ListenUDP(ctx context.Context, address net.Address, port net.Port, streamSe
|
||||
}
|
||||
|
||||
var err error
|
||||
hub.conn, err = internet.ListenSystemPacket(ctx, &net.UDPAddr{
|
||||
IP: address.IP(),
|
||||
Port: int(port),
|
||||
}, sockopt)
|
||||
if streamSettings.FinalMask != nil {
|
||||
hub.conn, err = streamSettings.FinalMask.ListenPacket(ctx, &net.UDPAddr{IP: address.IP(), Port: int(port)})
|
||||
} else {
|
||||
hub.conn, err = internet.ListenSystemPacket(ctx, &net.UDPAddr{IP: address.IP(), Port: int(port)}, streamSettings.SocketSettings)
|
||||
}
|
||||
if err != nil {
|
||||
return nil, err
|
||||
}
|
||||
|
||||
raw := hub.conn
|
||||
|
||||
if streamSettings.UdpmaskManager != nil {
|
||||
hub.conn, err = streamSettings.UdpmaskManager.WrapPacketConnServer(raw)
|
||||
if err != nil {
|
||||
raw.Close()
|
||||
return nil, errors.New("mask err").Base(err)
|
||||
}
|
||||
}
|
||||
|
||||
errors.LogInfo(ctx, "listening UDP on ", address, ":", port)
|
||||
hub.udpConn, _ = hub.conn.(*net.UDPConn)
|
||||
hub.cache = make(chan *udp.Packet, hub.capacity)
|
||||
|
||||
@@ -48,20 +48,16 @@ func dialWebSocket(ctx context.Context, dest net.Destination, streamSettings *in
|
||||
|
||||
dialer := &websocket.Dialer{
|
||||
NetDial: func(network, addr string) (net.Conn, error) {
|
||||
conn, err := internet.DialSystem(ctx, dest, streamSettings.SocketSettings)
|
||||
var conn net.Conn
|
||||
var err error
|
||||
if streamSettings.FinalMask != nil {
|
||||
conn, err = streamSettings.FinalMask.DialTCP(ctx, dest)
|
||||
} else {
|
||||
conn, err = internet.DialSystem(ctx, dest, streamSettings.SocketSettings)
|
||||
}
|
||||
if err != nil {
|
||||
return nil, err
|
||||
return nil, errors.New("failed to dial to dest").Base(err)
|
||||
}
|
||||
|
||||
if streamSettings.TcpmaskManager != nil {
|
||||
newConn, err := streamSettings.TcpmaskManager.WrapConnClient(conn)
|
||||
if err != nil {
|
||||
conn.Close()
|
||||
return nil, errors.New("mask err").Base(err)
|
||||
}
|
||||
conn = newConn
|
||||
}
|
||||
|
||||
return conn, err
|
||||
},
|
||||
ReadBufferSize: 4 * 1024,
|
||||
@@ -79,19 +75,15 @@ func dialWebSocket(ctx context.Context, dest net.Destination, streamSettings *in
|
||||
if fingerprint := tls.GetFingerprint(tConfig.Fingerprint); fingerprint != nil {
|
||||
dialer.NetDialTLSContext = func(_ context.Context, _, addr string) (net.Conn, error) {
|
||||
// Like the NetDial in the dialer
|
||||
pconn, err := internet.DialSystem(ctx, dest, streamSettings.SocketSettings)
|
||||
if err != nil {
|
||||
errors.LogErrorInner(ctx, err, "failed to dial to "+addr)
|
||||
return nil, err
|
||||
var pconn net.Conn
|
||||
var err error
|
||||
if streamSettings.FinalMask != nil {
|
||||
pconn, err = streamSettings.FinalMask.DialTCP(ctx, dest)
|
||||
} else {
|
||||
pconn, err = internet.DialSystem(ctx, dest, streamSettings.SocketSettings)
|
||||
}
|
||||
|
||||
if streamSettings.TcpmaskManager != nil {
|
||||
newConn, err := streamSettings.TcpmaskManager.WrapConnClient(pconn)
|
||||
if err != nil {
|
||||
pconn.Close()
|
||||
return nil, errors.New("mask err").Base(err)
|
||||
}
|
||||
pconn = newConn
|
||||
if err != nil {
|
||||
return nil, errors.New("failed to dial to dest").Base(err)
|
||||
}
|
||||
|
||||
// TLS and apply the handshake
|
||||
|
||||
@@ -97,29 +97,21 @@ func ListenWS(ctx context.Context, address net.Address, port net.Port, streamSet
|
||||
}
|
||||
var listener net.Listener
|
||||
var err error
|
||||
var addr net.Addr
|
||||
if port == net.Port(0) { // unix
|
||||
listener, err = internet.ListenSystem(ctx, &net.UnixAddr{
|
||||
Name: address.Domain(),
|
||||
Net: "unix",
|
||||
}, streamSettings.SocketSettings)
|
||||
if err != nil {
|
||||
return nil, errors.New("failed to listen unix domain socket(for WS) on ", address).Base(err)
|
||||
}
|
||||
errors.LogInfo(ctx, "listening unix domain socket(for WS) on ", address)
|
||||
addr = &net.UnixAddr{Name: address.Domain(), Net: "unix"}
|
||||
} else { // tcp
|
||||
listener, err = internet.ListenSystem(ctx, &net.TCPAddr{
|
||||
IP: address.IP(),
|
||||
Port: int(port),
|
||||
}, streamSettings.SocketSettings)
|
||||
if err != nil {
|
||||
return nil, errors.New("failed to listen TCP(for WS) on ", address, ":", port).Base(err)
|
||||
}
|
||||
errors.LogInfo(ctx, "listening TCP(for WS) on ", address, ":", port)
|
||||
addr = &net.TCPAddr{IP: address.IP(), Port: int(port)}
|
||||
}
|
||||
|
||||
if streamSettings.TcpmaskManager != nil {
|
||||
listener, _ = streamSettings.TcpmaskManager.WrapListener(listener)
|
||||
if streamSettings.FinalMask != nil {
|
||||
listener, err = streamSettings.FinalMask.Listen(ctx, addr)
|
||||
} else {
|
||||
listener, err = internet.ListenSystem(ctx, addr, streamSettings.SocketSettings)
|
||||
}
|
||||
if err != nil {
|
||||
return nil, errors.New("failed to listen ", addr.Network(), "(for WS) on ", address, ":", port).Base(err)
|
||||
}
|
||||
errors.LogInfo(ctx, "listening ", addr.Network(), "(for WS) on ", address, ":", port)
|
||||
|
||||
if streamSettings.SocketSettings != nil && streamSettings.SocketSettings.AcceptProxyProtocol {
|
||||
errors.LogWarning(ctx, "accepting PROXY protocol")
|
||||
|
||||
@@ -0,0 +1,153 @@
|
||||
package xdrive
|
||||
|
||||
import (
|
||||
"context"
|
||||
gotls "crypto/tls"
|
||||
"net/http"
|
||||
"time"
|
||||
|
||||
"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/reality"
|
||||
"github.com/xtls/xray-core/transport/internet/tls"
|
||||
"golang.org/x/net/http2"
|
||||
)
|
||||
|
||||
type serviceTransport struct {
|
||||
plain http.RoundTripper
|
||||
secure http.RoundTripper
|
||||
}
|
||||
|
||||
func (t *serviceTransport) RoundTrip(r *http.Request) (*http.Response, error) {
|
||||
if r.URL.Scheme == "https" {
|
||||
return t.secure.RoundTrip(r)
|
||||
}
|
||||
return t.plain.RoundTrip(r)
|
||||
}
|
||||
|
||||
func newServiceClient(streamSettings *internet.MemoryStreamConfig, timeout time.Duration, maxConns int) *http.Client {
|
||||
var (
|
||||
tlsConfig *tls.Config
|
||||
realityConfig *reality.Config
|
||||
sockopt *internet.SocketConfig
|
||||
fronting *net.Destination
|
||||
)
|
||||
if streamSettings != nil {
|
||||
tlsConfig = tls.ConfigFromStreamSettings(streamSettings)
|
||||
realityConfig = reality.ConfigFromStreamSettings(streamSettings)
|
||||
sockopt = streamSettings.SocketSettings
|
||||
fronting = streamSettings.Destination
|
||||
}
|
||||
overHTTP2 := allowsHTTP2(tlsConfig, realityConfig)
|
||||
|
||||
dial := func(ctx context.Context, addr string) (net.Conn, net.Destination, error) {
|
||||
host, err := net.ParseDestination("tcp:" + addr)
|
||||
if err != nil {
|
||||
return nil, host, errors.New("bad address: ", addr).Base(err)
|
||||
}
|
||||
|
||||
target := host
|
||||
if fronting != nil {
|
||||
target.Address = fronting.Address
|
||||
if fronting.Port != 0 {
|
||||
target.Port = fronting.Port
|
||||
}
|
||||
}
|
||||
|
||||
var conn net.Conn
|
||||
if streamSettings.FinalMask != nil {
|
||||
conn, err = streamSettings.FinalMask.DialTCP(ctx, target)
|
||||
} else {
|
||||
conn, err = internet.DialSystem(ctx, target, sockopt)
|
||||
}
|
||||
if err != nil {
|
||||
return nil, host, errors.New("failed to dial to dest").Base(err)
|
||||
}
|
||||
return conn, host, nil
|
||||
}
|
||||
|
||||
dialPlain := func(ctx context.Context, network, addr string) (net.Conn, error) {
|
||||
conn, _, err := dial(ctx, addr)
|
||||
return conn, err
|
||||
}
|
||||
|
||||
dialTLS := func(ctx context.Context, addr string) (net.Conn, error) {
|
||||
conn, host, err := dial(ctx, addr)
|
||||
if err != nil {
|
||||
return nil, err
|
||||
}
|
||||
|
||||
if realityConfig != nil {
|
||||
return reality.UClient(conn, realityConfig, ctx, host)
|
||||
}
|
||||
|
||||
gotlsConfig := &gotls.Config{ServerName: host.Address.String()}
|
||||
if tlsConfig != nil {
|
||||
gotlsConfig = tlsConfig.GetTLSConfig(tls.WithDestination(host))
|
||||
}
|
||||
if len(gotlsConfig.NextProtos) != 1 {
|
||||
if overHTTP2 {
|
||||
gotlsConfig.NextProtos = []string{"h2"}
|
||||
} else {
|
||||
gotlsConfig.NextProtos = []string{"http/1.1"}
|
||||
}
|
||||
}
|
||||
|
||||
if tlsConfig != nil {
|
||||
if fingerprint := tls.GetFingerprint(tlsConfig.Fingerprint); fingerprint != nil {
|
||||
uconn := tls.UClient(conn, gotlsConfig, fingerprint)
|
||||
if err := uconn.(*tls.UConn).HandshakeContext(ctx); err != nil {
|
||||
conn.Close()
|
||||
return nil, err
|
||||
}
|
||||
return uconn, nil
|
||||
}
|
||||
}
|
||||
return tls.Client(conn, gotlsConfig), nil
|
||||
}
|
||||
|
||||
var secure http.RoundTripper
|
||||
if overHTTP2 {
|
||||
secure = &http2.Transport{
|
||||
DialTLSContext: func(ctx context.Context, network, addr string, cfg *gotls.Config) (net.Conn, error) {
|
||||
return dialTLS(ctx, addr)
|
||||
},
|
||||
IdleConnTimeout: net.ConnIdleTimeout,
|
||||
}
|
||||
} else {
|
||||
secure = &http.Transport{
|
||||
DialTLSContext: func(ctx context.Context, network, addr string) (net.Conn, error) {
|
||||
return dialTLS(ctx, addr)
|
||||
},
|
||||
IdleConnTimeout: net.ConnIdleTimeout,
|
||||
MaxIdleConns: maxConns,
|
||||
MaxIdleConnsPerHost: maxConns,
|
||||
MaxConnsPerHost: maxConns,
|
||||
}
|
||||
}
|
||||
|
||||
return &http.Client{
|
||||
Transport: &serviceTransport{
|
||||
plain: &http.Transport{
|
||||
DialContext: dialPlain,
|
||||
IdleConnTimeout: net.ConnIdleTimeout,
|
||||
MaxIdleConns: maxConns,
|
||||
MaxIdleConnsPerHost: maxConns,
|
||||
MaxConnsPerHost: maxConns,
|
||||
},
|
||||
secure: secure,
|
||||
},
|
||||
Timeout: timeout,
|
||||
}
|
||||
}
|
||||
|
||||
func allowsHTTP2(tlsConfig *tls.Config, realityConfig *reality.Config) bool {
|
||||
if realityConfig != nil {
|
||||
return true
|
||||
}
|
||||
if tlsConfig == nil {
|
||||
return false
|
||||
}
|
||||
return len(tlsConfig.NextProtocol) == 1 && tlsConfig.NextProtocol[0] == "h2"
|
||||
}
|
||||
@@ -0,0 +1,101 @@
|
||||
package xdrive
|
||||
|
||||
import (
|
||||
"context"
|
||||
gotls "crypto/tls"
|
||||
"net"
|
||||
"net/http"
|
||||
"sync"
|
||||
"testing"
|
||||
"time"
|
||||
|
||||
xnet "github.com/xtls/xray-core/common/net"
|
||||
"github.com/xtls/xray-core/transport/internet"
|
||||
"github.com/xtls/xray-core/transport/internet/tls"
|
||||
)
|
||||
|
||||
func recordingTLSListener(t *testing.T, sni *string, mu *sync.Mutex) net.Listener {
|
||||
t.Helper()
|
||||
|
||||
ln, err := net.Listen("tcp", "127.0.0.1:0")
|
||||
if err != nil {
|
||||
t.Fatalf("listen: %v", err)
|
||||
}
|
||||
go func() {
|
||||
for {
|
||||
conn, err := ln.Accept()
|
||||
if err != nil {
|
||||
return
|
||||
}
|
||||
cfg := &gotls.Config{
|
||||
GetConfigForClient: func(hello *gotls.ClientHelloInfo) (*gotls.Config, error) {
|
||||
mu.Lock()
|
||||
*sni = hello.ServerName
|
||||
mu.Unlock()
|
||||
return nil, nil
|
||||
},
|
||||
}
|
||||
tconn := gotls.Server(conn, cfg)
|
||||
tconn.HandshakeContext(context.Background())
|
||||
tconn.Close()
|
||||
}
|
||||
}()
|
||||
t.Cleanup(func() { ln.Close() })
|
||||
return ln
|
||||
}
|
||||
|
||||
func sniForSettings(t *testing.T, serverName string) string {
|
||||
t.Helper()
|
||||
|
||||
var (
|
||||
sni string
|
||||
mu sync.Mutex
|
||||
)
|
||||
ln := recordingTLSListener(t, &sni, &mu)
|
||||
addr := ln.Addr().(*net.TCPAddr)
|
||||
|
||||
settings := &internet.MemoryStreamConfig{
|
||||
ProtocolName: protocolName,
|
||||
Destination: &xnet.Destination{
|
||||
Address: xnet.ParseAddress(addr.IP.String()),
|
||||
Port: xnet.Port(addr.Port),
|
||||
Network: xnet.Network_TCP,
|
||||
},
|
||||
SecuritySettings: &tls.Config{ServerName: serverName},
|
||||
}
|
||||
|
||||
prev := driveFilesURL
|
||||
driveFilesURL = "https://www.googleapis.com/drive/v3/files"
|
||||
defer func() { driveFilesURL = prev }()
|
||||
|
||||
client := newServiceClient(settings, 5*time.Second, 8)
|
||||
req, err := http.NewRequest(http.MethodGet, driveFilesURL, nil)
|
||||
if err != nil {
|
||||
t.Fatalf("request: %v", err)
|
||||
}
|
||||
client.Do(req)
|
||||
|
||||
deadline := time.Now().Add(3 * time.Second)
|
||||
for time.Now().Before(deadline) {
|
||||
mu.Lock()
|
||||
got := sni
|
||||
mu.Unlock()
|
||||
if got != "" {
|
||||
return got
|
||||
}
|
||||
time.Sleep(5 * time.Millisecond)
|
||||
}
|
||||
return ""
|
||||
}
|
||||
|
||||
func TestServiceSNIDefaultsToHost(t *testing.T) {
|
||||
if got := sniForSettings(t, ""); got != "www.googleapis.com" {
|
||||
t.Fatalf("SNI defaulted to %q, want the host www.googleapis.com, not address", got)
|
||||
}
|
||||
}
|
||||
|
||||
func TestServiceSNIOverride(t *testing.T) {
|
||||
if got := sniForSettings(t, "www.google.com"); got != "www.google.com" {
|
||||
t.Fatalf("explicit serverName gave SNI %q, want www.google.com", got)
|
||||
}
|
||||
}
|
||||
@@ -0,0 +1,222 @@
|
||||
// Code generated by protoc-gen-go. DO NOT EDIT.
|
||||
// versions:
|
||||
// protoc-gen-go v1.36.11
|
||||
// protoc v6.33.5
|
||||
// source: transport/internet/xdrive/config.proto
|
||||
|
||||
package xdrive
|
||||
|
||||
import (
|
||||
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 Config struct {
|
||||
state protoimpl.MessageState `protogen:"open.v1"`
|
||||
RemoteFolder string `protobuf:"bytes,1,opt,name=remote_folder,json=remoteFolder,proto3" json:"remote_folder,omitempty"`
|
||||
Service string `protobuf:"bytes,2,opt,name=service,proto3" json:"service,omitempty"`
|
||||
Secrets []string `protobuf:"bytes,3,rep,name=secrets,proto3" json:"secrets,omitempty"`
|
||||
SegmentBytes uint32 `protobuf:"varint,4,opt,name=segment_bytes,json=segmentBytes,proto3" json:"segment_bytes,omitempty"`
|
||||
FlushIntervalMs uint32 `protobuf:"varint,5,opt,name=flush_interval_ms,json=flushIntervalMs,proto3" json:"flush_interval_ms,omitempty"`
|
||||
PollIntervalMs uint32 `protobuf:"varint,6,opt,name=poll_interval_ms,json=pollIntervalMs,proto3" json:"poll_interval_ms,omitempty"`
|
||||
MaxPollIntervalMs uint32 `protobuf:"varint,7,opt,name=max_poll_interval_ms,json=maxPollIntervalMs,proto3" json:"max_poll_interval_ms,omitempty"`
|
||||
SessionTtlSeconds uint32 `protobuf:"varint,8,opt,name=session_ttl_seconds,json=sessionTtlSeconds,proto3" json:"session_ttl_seconds,omitempty"`
|
||||
Concurrency uint32 `protobuf:"varint,9,opt,name=concurrency,proto3" json:"concurrency,omitempty"`
|
||||
EagerWindowMs uint32 `protobuf:"varint,10,opt,name=eager_window_ms,json=eagerWindowMs,proto3" json:"eager_window_ms,omitempty"`
|
||||
HoleTimeoutMs uint32 `protobuf:"varint,11,opt,name=hole_timeout_ms,json=holeTimeoutMs,proto3" json:"hole_timeout_ms,omitempty"`
|
||||
Template string `protobuf:"bytes,12,opt,name=template,proto3" json:"template,omitempty"`
|
||||
unknownFields protoimpl.UnknownFields
|
||||
sizeCache protoimpl.SizeCache
|
||||
}
|
||||
|
||||
func (x *Config) Reset() {
|
||||
*x = Config{}
|
||||
mi := &file_transport_internet_xdrive_config_proto_msgTypes[0]
|
||||
ms := protoimpl.X.MessageStateOf(protoimpl.Pointer(x))
|
||||
ms.StoreMessageInfo(mi)
|
||||
}
|
||||
|
||||
func (x *Config) String() string {
|
||||
return protoimpl.X.MessageStringOf(x)
|
||||
}
|
||||
|
||||
func (*Config) ProtoMessage() {}
|
||||
|
||||
func (x *Config) ProtoReflect() protoreflect.Message {
|
||||
mi := &file_transport_internet_xdrive_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 Config.ProtoReflect.Descriptor instead.
|
||||
func (*Config) Descriptor() ([]byte, []int) {
|
||||
return file_transport_internet_xdrive_config_proto_rawDescGZIP(), []int{0}
|
||||
}
|
||||
|
||||
func (x *Config) GetRemoteFolder() string {
|
||||
if x != nil {
|
||||
return x.RemoteFolder
|
||||
}
|
||||
return ""
|
||||
}
|
||||
|
||||
func (x *Config) GetService() string {
|
||||
if x != nil {
|
||||
return x.Service
|
||||
}
|
||||
return ""
|
||||
}
|
||||
|
||||
func (x *Config) GetSecrets() []string {
|
||||
if x != nil {
|
||||
return x.Secrets
|
||||
}
|
||||
return nil
|
||||
}
|
||||
|
||||
func (x *Config) GetSegmentBytes() uint32 {
|
||||
if x != nil {
|
||||
return x.SegmentBytes
|
||||
}
|
||||
return 0
|
||||
}
|
||||
|
||||
func (x *Config) GetFlushIntervalMs() uint32 {
|
||||
if x != nil {
|
||||
return x.FlushIntervalMs
|
||||
}
|
||||
return 0
|
||||
}
|
||||
|
||||
func (x *Config) GetPollIntervalMs() uint32 {
|
||||
if x != nil {
|
||||
return x.PollIntervalMs
|
||||
}
|
||||
return 0
|
||||
}
|
||||
|
||||
func (x *Config) GetMaxPollIntervalMs() uint32 {
|
||||
if x != nil {
|
||||
return x.MaxPollIntervalMs
|
||||
}
|
||||
return 0
|
||||
}
|
||||
|
||||
func (x *Config) GetSessionTtlSeconds() uint32 {
|
||||
if x != nil {
|
||||
return x.SessionTtlSeconds
|
||||
}
|
||||
return 0
|
||||
}
|
||||
|
||||
func (x *Config) GetConcurrency() uint32 {
|
||||
if x != nil {
|
||||
return x.Concurrency
|
||||
}
|
||||
return 0
|
||||
}
|
||||
|
||||
func (x *Config) GetEagerWindowMs() uint32 {
|
||||
if x != nil {
|
||||
return x.EagerWindowMs
|
||||
}
|
||||
return 0
|
||||
}
|
||||
|
||||
func (x *Config) GetHoleTimeoutMs() uint32 {
|
||||
if x != nil {
|
||||
return x.HoleTimeoutMs
|
||||
}
|
||||
return 0
|
||||
}
|
||||
|
||||
func (x *Config) GetTemplate() string {
|
||||
if x != nil {
|
||||
return x.Template
|
||||
}
|
||||
return ""
|
||||
}
|
||||
|
||||
var File_transport_internet_xdrive_config_proto protoreflect.FileDescriptor
|
||||
|
||||
const file_transport_internet_xdrive_config_proto_rawDesc = "" +
|
||||
"\n" +
|
||||
"&transport/internet/xdrive/config.proto\x12\x1exray.transport.internet.xdrive\"\xcb\x03\n" +
|
||||
"\x06Config\x12#\n" +
|
||||
"\rremote_folder\x18\x01 \x01(\tR\fremoteFolder\x12\x18\n" +
|
||||
"\aservice\x18\x02 \x01(\tR\aservice\x12\x18\n" +
|
||||
"\asecrets\x18\x03 \x03(\tR\asecrets\x12#\n" +
|
||||
"\rsegment_bytes\x18\x04 \x01(\rR\fsegmentBytes\x12*\n" +
|
||||
"\x11flush_interval_ms\x18\x05 \x01(\rR\x0fflushIntervalMs\x12(\n" +
|
||||
"\x10poll_interval_ms\x18\x06 \x01(\rR\x0epollIntervalMs\x12/\n" +
|
||||
"\x14max_poll_interval_ms\x18\a \x01(\rR\x11maxPollIntervalMs\x12.\n" +
|
||||
"\x13session_ttl_seconds\x18\b \x01(\rR\x11sessionTtlSeconds\x12 \n" +
|
||||
"\vconcurrency\x18\t \x01(\rR\vconcurrency\x12&\n" +
|
||||
"\x0feager_window_ms\x18\n" +
|
||||
" \x01(\rR\reagerWindowMs\x12&\n" +
|
||||
"\x0fhole_timeout_ms\x18\v \x01(\rR\rholeTimeoutMs\x12\x1a\n" +
|
||||
"\btemplate\x18\f \x01(\tR\btemplateB5Z3github.com/xtls/xray-core/transport/internet/xdriveb\x06proto3"
|
||||
|
||||
var (
|
||||
file_transport_internet_xdrive_config_proto_rawDescOnce sync.Once
|
||||
file_transport_internet_xdrive_config_proto_rawDescData []byte
|
||||
)
|
||||
|
||||
func file_transport_internet_xdrive_config_proto_rawDescGZIP() []byte {
|
||||
file_transport_internet_xdrive_config_proto_rawDescOnce.Do(func() {
|
||||
file_transport_internet_xdrive_config_proto_rawDescData = protoimpl.X.CompressGZIP(unsafe.Slice(unsafe.StringData(file_transport_internet_xdrive_config_proto_rawDesc), len(file_transport_internet_xdrive_config_proto_rawDesc)))
|
||||
})
|
||||
return file_transport_internet_xdrive_config_proto_rawDescData
|
||||
}
|
||||
|
||||
var file_transport_internet_xdrive_config_proto_msgTypes = make([]protoimpl.MessageInfo, 1)
|
||||
var file_transport_internet_xdrive_config_proto_goTypes = []any{
|
||||
(*Config)(nil), // 0: xray.transport.internet.xdrive.Config
|
||||
}
|
||||
var file_transport_internet_xdrive_config_proto_depIdxs = []int32{
|
||||
0, // [0:0] is the sub-list for method output_type
|
||||
0, // [0:0] is the sub-list for method input_type
|
||||
0, // [0:0] is the sub-list for extension type_name
|
||||
0, // [0:0] is the sub-list for extension extendee
|
||||
0, // [0:0] is the sub-list for field type_name
|
||||
}
|
||||
|
||||
func init() { file_transport_internet_xdrive_config_proto_init() }
|
||||
func file_transport_internet_xdrive_config_proto_init() {
|
||||
if File_transport_internet_xdrive_config_proto != nil {
|
||||
return
|
||||
}
|
||||
type x struct{}
|
||||
out := protoimpl.TypeBuilder{
|
||||
File: protoimpl.DescBuilder{
|
||||
GoPackagePath: reflect.TypeOf(x{}).PkgPath(),
|
||||
RawDescriptor: unsafe.Slice(unsafe.StringData(file_transport_internet_xdrive_config_proto_rawDesc), len(file_transport_internet_xdrive_config_proto_rawDesc)),
|
||||
NumEnums: 0,
|
||||
NumMessages: 1,
|
||||
NumExtensions: 0,
|
||||
NumServices: 0,
|
||||
},
|
||||
GoTypes: file_transport_internet_xdrive_config_proto_goTypes,
|
||||
DependencyIndexes: file_transport_internet_xdrive_config_proto_depIdxs,
|
||||
MessageInfos: file_transport_internet_xdrive_config_proto_msgTypes,
|
||||
}.Build()
|
||||
File_transport_internet_xdrive_config_proto = out.File
|
||||
file_transport_internet_xdrive_config_proto_goTypes = nil
|
||||
file_transport_internet_xdrive_config_proto_depIdxs = nil
|
||||
}
|
||||
@@ -0,0 +1,19 @@
|
||||
syntax = "proto3";
|
||||
|
||||
package xray.transport.internet.xdrive;
|
||||
option go_package = "github.com/xtls/xray-core/transport/internet/xdrive";
|
||||
|
||||
message Config {
|
||||
string remote_folder = 1;
|
||||
string service = 2;
|
||||
repeated string secrets = 3;
|
||||
uint32 segment_bytes = 4;
|
||||
uint32 flush_interval_ms = 5;
|
||||
uint32 poll_interval_ms = 6;
|
||||
uint32 max_poll_interval_ms = 7;
|
||||
uint32 session_ttl_seconds = 8;
|
||||
uint32 concurrency = 9;
|
||||
uint32 eager_window_ms = 10;
|
||||
uint32 hole_timeout_ms = 11;
|
||||
string template = 12;
|
||||
}
|
||||
@@ -0,0 +1,139 @@
|
||||
package xdrive
|
||||
|
||||
import (
|
||||
"context"
|
||||
"os"
|
||||
"sync"
|
||||
"time"
|
||||
|
||||
"github.com/xtls/xray-core/common/net"
|
||||
)
|
||||
|
||||
var placeholderAddr = &net.TCPAddr{IP: net.IP{127, 0, 0, 1}, Port: 0}
|
||||
|
||||
type Conn struct {
|
||||
cancel context.CancelFunc
|
||||
writer *walWriter
|
||||
reader *walReader
|
||||
onClose func()
|
||||
|
||||
readBuf []byte
|
||||
|
||||
deadlineMu sync.Mutex
|
||||
readDeadline time.Time
|
||||
writeDeadline time.Time
|
||||
|
||||
closeOnce sync.Once
|
||||
closeErr error
|
||||
}
|
||||
|
||||
func newConn(ctx context.Context, storage Storage, writePrefix, readPrefix string, p params, onClose func()) *Conn {
|
||||
ctx, cancel := context.WithCancel(ctx)
|
||||
return &Conn{
|
||||
cancel: cancel,
|
||||
writer: newWALWriter(ctx, storage, writePrefix, p),
|
||||
reader: newWALReader(ctx, storage, readPrefix, p),
|
||||
onClose: onClose,
|
||||
}
|
||||
}
|
||||
|
||||
func (c *Conn) Read(b []byte) (int, error) {
|
||||
if len(c.readBuf) == 0 {
|
||||
data, err := c.receive()
|
||||
if err != nil {
|
||||
return 0, err
|
||||
}
|
||||
c.readBuf = data
|
||||
}
|
||||
n := copy(b, c.readBuf)
|
||||
c.readBuf = c.readBuf[n:]
|
||||
return n, nil
|
||||
}
|
||||
|
||||
func (c *Conn) receive() ([]byte, error) {
|
||||
deadline := c.getDeadline(true)
|
||||
if deadline.IsZero() {
|
||||
data, ok := <-c.reader.ch
|
||||
if !ok {
|
||||
return nil, c.reader.Err()
|
||||
}
|
||||
return data, nil
|
||||
}
|
||||
|
||||
if !time.Now().Before(deadline) {
|
||||
return nil, os.ErrDeadlineExceeded
|
||||
}
|
||||
timer := time.NewTimer(time.Until(deadline))
|
||||
defer timer.Stop()
|
||||
|
||||
select {
|
||||
case data, ok := <-c.reader.ch:
|
||||
if !ok {
|
||||
return nil, c.reader.Err()
|
||||
}
|
||||
return data, nil
|
||||
case <-timer.C:
|
||||
return nil, os.ErrDeadlineExceeded
|
||||
}
|
||||
}
|
||||
|
||||
func (c *Conn) Write(b []byte) (int, error) {
|
||||
if deadline := c.getDeadline(false); !deadline.IsZero() && !time.Now().Before(deadline) {
|
||||
return 0, os.ErrDeadlineExceeded
|
||||
}
|
||||
n, err := c.writer.Write(b)
|
||||
if err == nil {
|
||||
c.reader.Wake()
|
||||
}
|
||||
return n, err
|
||||
}
|
||||
|
||||
func (c *Conn) Close() error {
|
||||
c.closeOnce.Do(func() {
|
||||
c.closeErr = c.writer.Close()
|
||||
c.cancel()
|
||||
if c.onClose != nil {
|
||||
c.onClose()
|
||||
}
|
||||
})
|
||||
return c.closeErr
|
||||
}
|
||||
|
||||
func (c *Conn) LocalAddr() net.Addr {
|
||||
return placeholderAddr
|
||||
}
|
||||
|
||||
func (c *Conn) RemoteAddr() net.Addr {
|
||||
return placeholderAddr
|
||||
}
|
||||
|
||||
func (c *Conn) getDeadline(read bool) time.Time {
|
||||
c.deadlineMu.Lock()
|
||||
defer c.deadlineMu.Unlock()
|
||||
if read {
|
||||
return c.readDeadline
|
||||
}
|
||||
return c.writeDeadline
|
||||
}
|
||||
|
||||
func (c *Conn) SetDeadline(t time.Time) error {
|
||||
c.deadlineMu.Lock()
|
||||
defer c.deadlineMu.Unlock()
|
||||
c.readDeadline = t
|
||||
c.writeDeadline = t
|
||||
return nil
|
||||
}
|
||||
|
||||
func (c *Conn) SetReadDeadline(t time.Time) error {
|
||||
c.deadlineMu.Lock()
|
||||
defer c.deadlineMu.Unlock()
|
||||
c.readDeadline = t
|
||||
return nil
|
||||
}
|
||||
|
||||
func (c *Conn) SetWriteDeadline(t time.Time) error {
|
||||
c.deadlineMu.Lock()
|
||||
defer c.deadlineMu.Unlock()
|
||||
c.writeDeadline = t
|
||||
return nil
|
||||
}
|
||||
@@ -0,0 +1,575 @@
|
||||
package xdrive
|
||||
|
||||
import (
|
||||
"bytes"
|
||||
"context"
|
||||
"encoding/base64"
|
||||
"encoding/json"
|
||||
"fmt"
|
||||
"io"
|
||||
"net/http"
|
||||
"net/url"
|
||||
"strings"
|
||||
"sync"
|
||||
"time"
|
||||
|
||||
"github.com/xtls/xray-core/common/dice"
|
||||
"github.com/xtls/xray-core/common/errors"
|
||||
"github.com/xtls/xray-core/transport/internet"
|
||||
)
|
||||
|
||||
const (
|
||||
flatSeparator = "~"
|
||||
|
||||
driveBoundary = "xdrive-boundary"
|
||||
drivePageSize = 1000
|
||||
driveMaxAttempts = 8
|
||||
driveMaxInflight = 32
|
||||
driveInlineLimit = 12000
|
||||
driveTimeout = 60 * time.Second
|
||||
)
|
||||
|
||||
var (
|
||||
driveMaxBackoff = 8 * time.Second
|
||||
driveTokenURL = "https://oauth2.googleapis.com/token"
|
||||
driveFilesURL = "https://www.googleapis.com/drive/v3/files"
|
||||
driveUploadURL = "https://www.googleapis.com/upload/drive/v3/files?uploadType=multipart&fields=id,name"
|
||||
driveInitialBackoff = 200 * time.Millisecond
|
||||
)
|
||||
|
||||
type driveStorage struct {
|
||||
folder string
|
||||
clientID string
|
||||
clientSecret string
|
||||
refreshToken string
|
||||
client *http.Client
|
||||
tokenURL string
|
||||
filesURL string
|
||||
uploadURL string
|
||||
backoff time.Duration
|
||||
|
||||
tokenMu sync.Mutex
|
||||
token string
|
||||
tokenExpiry time.Time
|
||||
|
||||
inflight chan struct{}
|
||||
|
||||
idMu sync.Mutex
|
||||
ids map[string]string
|
||||
}
|
||||
|
||||
func newDriveStorage(streamSettings *internet.MemoryStreamConfig, config *Config) (*driveStorage, error) {
|
||||
if config.RemoteFolder == "" {
|
||||
return nil, errors.New(`empty "remoteFolder", it must be a Google Drive folder id`)
|
||||
}
|
||||
if len(config.Secrets) != 3 {
|
||||
return nil, errors.New("Google Drive needs 3 secrets in order of ClientID, ClientSecret, RefreshToken")
|
||||
}
|
||||
for i, secret := range config.Secrets {
|
||||
if secret == "" {
|
||||
return nil, errors.New("Google Drive secret ", i, " is empty")
|
||||
}
|
||||
}
|
||||
|
||||
return &driveStorage{
|
||||
folder: config.RemoteFolder,
|
||||
clientID: config.Secrets[0],
|
||||
clientSecret: config.Secrets[1],
|
||||
refreshToken: config.Secrets[2],
|
||||
client: newServiceClient(streamSettings, driveTimeout, driveMaxInflight),
|
||||
tokenURL: driveTokenURL,
|
||||
filesURL: driveFilesURL,
|
||||
uploadURL: driveUploadURL,
|
||||
backoff: driveInitialBackoff,
|
||||
inflight: make(chan struct{}, driveMaxInflight),
|
||||
ids: make(map[string]string),
|
||||
}, nil
|
||||
}
|
||||
|
||||
func flatten(name string) string {
|
||||
return strings.ReplaceAll(name, "/", flatSeparator)
|
||||
}
|
||||
|
||||
func quoteDriveValue(value string) string {
|
||||
return strings.NewReplacer(`\`, `\\`, `'`, `\'`).Replace(value)
|
||||
}
|
||||
|
||||
func (s *driveStorage) accessToken(ctx context.Context) (string, error) {
|
||||
s.tokenMu.Lock()
|
||||
defer s.tokenMu.Unlock()
|
||||
|
||||
if s.token != "" && time.Now().Before(s.tokenExpiry) {
|
||||
return s.token, nil
|
||||
}
|
||||
|
||||
form := url.Values{
|
||||
"client_id": {s.clientID},
|
||||
"client_secret": {s.clientSecret},
|
||||
"refresh_token": {s.refreshToken},
|
||||
"grant_type": {"refresh_token"},
|
||||
}
|
||||
req, err := http.NewRequestWithContext(ctx, http.MethodPost, s.tokenURL, strings.NewReader(form.Encode()))
|
||||
if err != nil {
|
||||
return "", errors.New("failed to build the token request").Base(err)
|
||||
}
|
||||
req.Header.Set("Content-Type", "application/x-www-form-urlencoded")
|
||||
|
||||
resp, err := s.client.Do(req)
|
||||
if err != nil {
|
||||
return "", errors.New("failed to refresh the access token").Base(err)
|
||||
}
|
||||
defer resp.Body.Close()
|
||||
|
||||
body, err := io.ReadAll(io.LimitReader(resp.Body, 1<<20))
|
||||
if err != nil {
|
||||
return "", errors.New("failed to read the token response").Base(err)
|
||||
}
|
||||
if resp.StatusCode != http.StatusOK {
|
||||
return "", errors.New("the token endpoint answered ", resp.StatusCode, ": ", string(body))
|
||||
}
|
||||
|
||||
var parsed struct {
|
||||
AccessToken string `json:"access_token"`
|
||||
ExpiresIn int64 `json:"expires_in"`
|
||||
}
|
||||
if err := json.Unmarshal(body, &parsed); err != nil {
|
||||
return "", errors.New("failed to parse the token response").Base(err)
|
||||
}
|
||||
if parsed.AccessToken == "" {
|
||||
return "", errors.New("the token endpoint returned no access token")
|
||||
}
|
||||
|
||||
lifetime := parsed.ExpiresIn
|
||||
if lifetime > 60 {
|
||||
lifetime -= 60
|
||||
}
|
||||
s.token = parsed.AccessToken
|
||||
s.tokenExpiry = time.Now().Add(time.Duration(lifetime) * time.Second)
|
||||
return s.token, nil
|
||||
}
|
||||
|
||||
func jitter(backoff time.Duration) time.Duration {
|
||||
half := backoff / 2
|
||||
if half <= 0 {
|
||||
return backoff
|
||||
}
|
||||
return half + time.Duration(dice.Roll(int(half)))
|
||||
}
|
||||
|
||||
func rateLimited(payload []byte) bool {
|
||||
var parsed struct {
|
||||
Error struct {
|
||||
Status string `json:"status"`
|
||||
Errors []struct {
|
||||
Reason string `json:"reason"`
|
||||
} `json:"errors"`
|
||||
} `json:"error"`
|
||||
}
|
||||
if json.Unmarshal(payload, &parsed) != nil {
|
||||
return false
|
||||
}
|
||||
for _, item := range parsed.Error.Errors {
|
||||
switch item.Reason {
|
||||
case "rateLimitExceeded", "userRateLimitExceeded", "sharingRateLimitExceeded":
|
||||
return true
|
||||
}
|
||||
}
|
||||
return parsed.Error.Status == "RESOURCE_EXHAUSTED"
|
||||
}
|
||||
|
||||
func retryableStatus(status int) bool {
|
||||
switch status {
|
||||
case http.StatusTooManyRequests, http.StatusInternalServerError,
|
||||
http.StatusBadGateway, http.StatusServiceUnavailable, http.StatusGatewayTimeout:
|
||||
return true
|
||||
}
|
||||
return false
|
||||
}
|
||||
|
||||
func (s *driveStorage) do(ctx context.Context, method, target, contentType string, body []byte) (int, []byte, error) {
|
||||
backoff := s.backoff
|
||||
var lastErr error
|
||||
|
||||
for attempt := 0; attempt < driveMaxAttempts; attempt++ {
|
||||
if attempt > 0 {
|
||||
select {
|
||||
case <-ctx.Done():
|
||||
return 0, nil, ctx.Err()
|
||||
case <-time.After(jitter(backoff)):
|
||||
}
|
||||
backoff *= 2
|
||||
if backoff > driveMaxBackoff {
|
||||
backoff = driveMaxBackoff
|
||||
}
|
||||
}
|
||||
|
||||
token, err := s.accessToken(ctx)
|
||||
if err != nil {
|
||||
lastErr = err
|
||||
continue
|
||||
}
|
||||
|
||||
select {
|
||||
case s.inflight <- struct{}{}:
|
||||
case <-ctx.Done():
|
||||
return 0, nil, ctx.Err()
|
||||
}
|
||||
|
||||
var reader io.Reader
|
||||
if body != nil {
|
||||
reader = bytes.NewReader(body)
|
||||
}
|
||||
req, err := http.NewRequestWithContext(ctx, method, target, reader)
|
||||
if err != nil {
|
||||
return 0, nil, errors.New("failed to build a Drive request").Base(err)
|
||||
}
|
||||
req.Header.Set("Authorization", "Bearer "+token)
|
||||
if contentType != "" {
|
||||
req.Header.Set("Content-Type", contentType)
|
||||
}
|
||||
|
||||
resp, err := s.client.Do(req)
|
||||
<-s.inflight
|
||||
if err != nil {
|
||||
if ctx.Err() != nil {
|
||||
return 0, nil, ctx.Err()
|
||||
}
|
||||
lastErr = errors.New("Drive request failed").Base(err)
|
||||
errors.LogWarningInner(ctx, err, "retrying a failed Drive request, attempt ",
|
||||
attempt+1, " of ", driveMaxAttempts)
|
||||
continue
|
||||
}
|
||||
payload, err := io.ReadAll(resp.Body)
|
||||
resp.Body.Close()
|
||||
if err != nil {
|
||||
if ctx.Err() != nil {
|
||||
return 0, nil, ctx.Err()
|
||||
}
|
||||
lastErr = errors.New("failed to read the Drive response").Base(err)
|
||||
continue
|
||||
}
|
||||
|
||||
if resp.StatusCode == http.StatusUnauthorized {
|
||||
s.invalidateToken()
|
||||
lastErr = errors.New("Drive rejected the access token")
|
||||
continue
|
||||
}
|
||||
if resp.StatusCode == http.StatusForbidden && rateLimited(payload) {
|
||||
lastErr = errors.New("Drive is rate limiting: ", string(payload))
|
||||
errors.LogWarning(ctx, "rate limited by Drive, attempt ",
|
||||
attempt+1, " of ", driveMaxAttempts)
|
||||
continue
|
||||
}
|
||||
if retryableStatus(resp.StatusCode) {
|
||||
lastErr = errors.New("Drive answered ", resp.StatusCode, ": ", string(payload))
|
||||
errors.LogWarning(ctx, "retrying after Drive answered ", resp.StatusCode,
|
||||
", attempt ", attempt+1, " of ", driveMaxAttempts)
|
||||
continue
|
||||
}
|
||||
return resp.StatusCode, payload, nil
|
||||
}
|
||||
|
||||
return 0, nil, lastErr
|
||||
}
|
||||
|
||||
func (s *driveStorage) invalidateToken() {
|
||||
s.tokenMu.Lock()
|
||||
s.token = ""
|
||||
s.tokenMu.Unlock()
|
||||
}
|
||||
|
||||
func (s *driveStorage) rememberID(name, id string) {
|
||||
s.idMu.Lock()
|
||||
s.ids[name] = id
|
||||
s.idMu.Unlock()
|
||||
}
|
||||
|
||||
func (s *driveStorage) forgetID(name string) {
|
||||
s.idMu.Lock()
|
||||
delete(s.ids, name)
|
||||
s.idMu.Unlock()
|
||||
}
|
||||
|
||||
func (s *driveStorage) cachedID(name string) (string, bool) {
|
||||
s.idMu.Lock()
|
||||
defer s.idMu.Unlock()
|
||||
id, ok := s.ids[name]
|
||||
return id, ok
|
||||
}
|
||||
|
||||
type driveFile struct {
|
||||
id string
|
||||
description string
|
||||
}
|
||||
|
||||
func (s *driveStorage) query(ctx context.Context, condition string) (map[string]driveFile, error) {
|
||||
found := make(map[string]driveFile)
|
||||
pageToken := ""
|
||||
|
||||
for {
|
||||
params := url.Values{
|
||||
"q": {"'" + quoteDriveValue(s.folder) + "' in parents and trashed = false and " + condition},
|
||||
"fields": {"nextPageToken,files(id,name,description)"},
|
||||
"pageSize": {fmt.Sprint(drivePageSize)},
|
||||
"supportsAllDrives": {"true"},
|
||||
"includeItemsFromAllDrives": {"true"},
|
||||
}
|
||||
if pageToken != "" {
|
||||
params.Set("pageToken", pageToken)
|
||||
}
|
||||
|
||||
status, payload, err := s.do(ctx, http.MethodGet, s.filesURL+"?"+params.Encode(), "", nil)
|
||||
if err != nil {
|
||||
return nil, err
|
||||
}
|
||||
if status != http.StatusOK {
|
||||
return nil, errors.New("Drive listing answered ", status, ": ", string(payload))
|
||||
}
|
||||
|
||||
var parsed struct {
|
||||
NextPageToken string `json:"nextPageToken"`
|
||||
Files []struct {
|
||||
ID string `json:"id"`
|
||||
Name string `json:"name"`
|
||||
Description string `json:"description"`
|
||||
} `json:"files"`
|
||||
}
|
||||
if err := json.Unmarshal(payload, &parsed); err != nil {
|
||||
return nil, errors.New("failed to parse the Drive listing").Base(err)
|
||||
}
|
||||
|
||||
for _, file := range parsed.Files {
|
||||
found[file.Name] = driveFile{id: file.ID, description: file.Description}
|
||||
s.rememberID(file.Name, file.ID)
|
||||
}
|
||||
|
||||
pageToken = parsed.NextPageToken
|
||||
if pageToken == "" {
|
||||
return found, nil
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
func (s *driveStorage) resolveID(ctx context.Context, flat string) (string, error) {
|
||||
if id, ok := s.cachedID(flat); ok {
|
||||
return id, nil
|
||||
}
|
||||
found, err := s.query(ctx, "name = '"+quoteDriveValue(flat)+"'")
|
||||
if err != nil {
|
||||
return "", err
|
||||
}
|
||||
if file, ok := found[flat]; ok {
|
||||
return file.id, nil
|
||||
}
|
||||
return "", errNotFound
|
||||
}
|
||||
|
||||
func (s *driveStorage) Put(ctx context.Context, name string, data []byte) error {
|
||||
if len(data) <= driveInlineLimit {
|
||||
return s.putInline(ctx, name, data)
|
||||
}
|
||||
return s.putMedia(ctx, name, data)
|
||||
}
|
||||
|
||||
func (s *driveStorage) putInline(ctx context.Context, name string, data []byte) error {
|
||||
flat := flatten(name)
|
||||
|
||||
body, err := json.Marshal(map[string]interface{}{
|
||||
"name": flat,
|
||||
"parents": []string{s.folder},
|
||||
"description": base64.StdEncoding.EncodeToString(data),
|
||||
})
|
||||
if err != nil {
|
||||
return errors.New("failed to build the inline metadata").Base(err)
|
||||
}
|
||||
|
||||
status, payload, err := s.do(ctx, http.MethodPost, s.filesURL+"?fields=id",
|
||||
"application/json; charset=UTF-8", body)
|
||||
if err != nil {
|
||||
return err
|
||||
}
|
||||
if status != http.StatusOK {
|
||||
return errors.New("Drive rejected the inline upload of ", name,
|
||||
" with ", status, ": ", string(payload))
|
||||
}
|
||||
|
||||
var parsed struct {
|
||||
ID string `json:"id"`
|
||||
}
|
||||
if err := json.Unmarshal(payload, &parsed); err != nil {
|
||||
return errors.New("failed to parse the Drive upload response").Base(err)
|
||||
}
|
||||
if parsed.ID != "" {
|
||||
s.rememberID(flat, parsed.ID)
|
||||
}
|
||||
return nil
|
||||
}
|
||||
|
||||
func (s *driveStorage) putMedia(ctx context.Context, name string, data []byte) error {
|
||||
flat := flatten(name)
|
||||
|
||||
metadata, err := json.Marshal(map[string]interface{}{
|
||||
"name": flat,
|
||||
"parents": []string{s.folder},
|
||||
})
|
||||
if err != nil {
|
||||
return errors.New("failed to build the upload metadata").Base(err)
|
||||
}
|
||||
|
||||
var body bytes.Buffer
|
||||
fmt.Fprintf(&body, "--%s\r\nContent-Type: application/json; charset=UTF-8\r\n\r\n", driveBoundary)
|
||||
body.Write(metadata)
|
||||
fmt.Fprintf(&body, "\r\n--%s\r\nContent-Type: application/octet-stream\r\n\r\n", driveBoundary)
|
||||
body.Write(data)
|
||||
fmt.Fprintf(&body, "\r\n--%s--\r\n", driveBoundary)
|
||||
|
||||
status, payload, err := s.do(ctx, http.MethodPost, s.uploadURL,
|
||||
"multipart/related; boundary="+driveBoundary, body.Bytes())
|
||||
if err != nil {
|
||||
return err
|
||||
}
|
||||
if status != http.StatusOK {
|
||||
return errors.New("Drive upload of ", name, " answered ", status, ": ", string(payload))
|
||||
}
|
||||
|
||||
var parsed struct {
|
||||
ID string `json:"id"`
|
||||
}
|
||||
if err := json.Unmarshal(payload, &parsed); err != nil {
|
||||
return errors.New("failed to parse the Drive upload response").Base(err)
|
||||
}
|
||||
if parsed.ID != "" {
|
||||
s.rememberID(flat, parsed.ID)
|
||||
}
|
||||
return nil
|
||||
}
|
||||
|
||||
func (s *driveStorage) Get(ctx context.Context, name string) ([]byte, error) {
|
||||
flat := flatten(name)
|
||||
id, err := s.resolveID(ctx, flat)
|
||||
if err != nil {
|
||||
return nil, err
|
||||
}
|
||||
|
||||
status, payload, err := s.do(ctx, http.MethodGet,
|
||||
s.filesURL+"/"+url.PathEscape(id)+"?alt=media&supportsAllDrives=true", "", nil)
|
||||
if err != nil {
|
||||
return nil, err
|
||||
}
|
||||
switch status {
|
||||
case http.StatusOK:
|
||||
if len(payload) > 0 {
|
||||
return payload, nil
|
||||
}
|
||||
return s.getInline(ctx, id)
|
||||
case http.StatusNotFound:
|
||||
s.forgetID(flat)
|
||||
return nil, errNotFound
|
||||
default:
|
||||
return nil, errors.New("Drive download of ", name, " answered ", status, ": ", string(payload))
|
||||
}
|
||||
}
|
||||
|
||||
func (s *driveStorage) getInline(ctx context.Context, id string) ([]byte, error) {
|
||||
status, payload, err := s.do(ctx, http.MethodGet,
|
||||
s.filesURL+"/"+url.PathEscape(id)+"?fields=description&supportsAllDrives=true", "", nil)
|
||||
if err != nil {
|
||||
return nil, err
|
||||
}
|
||||
if status != http.StatusOK {
|
||||
return nil, errors.New("Drive answered ", status, " for inline data: ", string(payload))
|
||||
}
|
||||
|
||||
var parsed struct {
|
||||
Description string `json:"description"`
|
||||
}
|
||||
if err := json.Unmarshal(payload, &parsed); err != nil {
|
||||
return nil, errors.New("failed to parse the inline data").Base(err)
|
||||
}
|
||||
if parsed.Description == "" {
|
||||
return nil, nil
|
||||
}
|
||||
data, err := base64.StdEncoding.DecodeString(parsed.Description)
|
||||
if err != nil {
|
||||
return nil, errors.New("the inline data is not valid base64").Base(err)
|
||||
}
|
||||
return data, nil
|
||||
}
|
||||
|
||||
func (s *driveStorage) deleteID(ctx context.Context, flat, id string) error {
|
||||
status, payload, err := s.do(ctx, http.MethodDelete,
|
||||
s.filesURL+"/"+url.PathEscape(id)+"?supportsAllDrives=true", "", nil)
|
||||
if err != nil {
|
||||
return err
|
||||
}
|
||||
s.forgetID(flat)
|
||||
switch status {
|
||||
case http.StatusOK, http.StatusNoContent, http.StatusNotFound:
|
||||
return nil
|
||||
default:
|
||||
return errors.New("Drive deletion of ", flat, " answered ", status, ": ", string(payload))
|
||||
}
|
||||
}
|
||||
|
||||
func (s *driveStorage) Delete(ctx context.Context, name string) error {
|
||||
flat := flatten(name)
|
||||
|
||||
if id, err := s.resolveID(ctx, flat); err == nil {
|
||||
if err := s.deleteID(ctx, flat, id); err != nil {
|
||||
return err
|
||||
}
|
||||
} else if err != errNotFound {
|
||||
return err
|
||||
}
|
||||
|
||||
children, err := s.query(ctx, "name contains '"+quoteDriveValue(flat+flatSeparator)+"'")
|
||||
if err != nil {
|
||||
return err
|
||||
}
|
||||
for childName, file := range children {
|
||||
if err := s.deleteID(ctx, childName, file.id); err != nil {
|
||||
return err
|
||||
}
|
||||
}
|
||||
return nil
|
||||
}
|
||||
|
||||
func (s *driveStorage) List(ctx context.Context, prefix string) ([]Entry, error) {
|
||||
flat := flatten(prefix) + flatSeparator
|
||||
|
||||
found, err := s.query(ctx, "name contains '"+quoteDriveValue(flat)+"'")
|
||||
if err != nil {
|
||||
return nil, err
|
||||
}
|
||||
|
||||
seen := make(map[string]bool, len(found))
|
||||
entries := make([]Entry, 0, len(found))
|
||||
for name, file := range found {
|
||||
rest := strings.TrimPrefix(name, flat)
|
||||
if rest == "" {
|
||||
continue
|
||||
}
|
||||
direct := true
|
||||
if cut := strings.Index(rest, flatSeparator); cut >= 0 {
|
||||
rest = rest[:cut]
|
||||
direct = false
|
||||
}
|
||||
if seen[rest] {
|
||||
continue
|
||||
}
|
||||
seen[rest] = true
|
||||
|
||||
entry := Entry{Name: rest}
|
||||
if direct && file.description != "" {
|
||||
if data, err := base64.StdEncoding.DecodeString(file.description); err == nil {
|
||||
entry.Inline = data
|
||||
}
|
||||
}
|
||||
entries = append(entries, entry)
|
||||
}
|
||||
return entries, nil
|
||||
}
|
||||
|
||||
func (s *driveStorage) Close() error {
|
||||
return nil
|
||||
}
|
||||
@@ -0,0 +1,345 @@
|
||||
package xdrive
|
||||
|
||||
import (
|
||||
"bytes"
|
||||
"context"
|
||||
"crypto/rand"
|
||||
"encoding/json"
|
||||
"fmt"
|
||||
"io"
|
||||
"os"
|
||||
"strconv"
|
||||
"sync"
|
||||
"testing"
|
||||
"time"
|
||||
|
||||
"github.com/xtls/xray-core/transport/internet"
|
||||
)
|
||||
|
||||
const liveSecretsEnv = "XRAY_XDRIVE_DRIVE_SECRETS"
|
||||
|
||||
func liveDriveConfig(t *testing.T) *Config {
|
||||
t.Helper()
|
||||
|
||||
path := os.Getenv(liveSecretsEnv)
|
||||
if path == "" {
|
||||
t.Skipf("set %s to run this test", liveSecretsEnv)
|
||||
}
|
||||
|
||||
payload, err := os.ReadFile(path)
|
||||
if err != nil {
|
||||
t.Fatalf("reading %s: %v", path, err)
|
||||
}
|
||||
var secrets struct {
|
||||
Folder string `json:"folder"`
|
||||
ClientID string `json:"client_id"`
|
||||
ClientSecret string `json:"client_secret"`
|
||||
RefreshToken string `json:"refresh_token"`
|
||||
}
|
||||
if err := json.Unmarshal(payload, &secrets); err != nil {
|
||||
t.Fatalf("parsing %s: %v", path, err)
|
||||
}
|
||||
|
||||
config := &Config{
|
||||
RemoteFolder: secrets.Folder,
|
||||
Service: "Google Drive",
|
||||
Secrets: []string{secrets.ClientID, secrets.ClientSecret, secrets.RefreshToken},
|
||||
SegmentBytes: 256 * 1024,
|
||||
FlushIntervalMs: 100,
|
||||
PollIntervalMs: 500,
|
||||
MaxPollIntervalMs: 2000,
|
||||
SessionTtlSeconds: 120,
|
||||
}
|
||||
if raw := os.Getenv("XRAY_XDRIVE_LIVE_SEGMENT"); raw != "" {
|
||||
config.SegmentBytes = uint32(envInt(t, "XRAY_XDRIVE_LIVE_SEGMENT"))
|
||||
}
|
||||
if raw := os.Getenv("XRAY_XDRIVE_LIVE_CONCURRENCY"); raw != "" {
|
||||
config.Concurrency = uint32(envInt(t, "XRAY_XDRIVE_LIVE_CONCURRENCY"))
|
||||
}
|
||||
|
||||
storage, err := newDriveStorage(nil, config)
|
||||
if err != nil {
|
||||
t.Fatalf("newDriveStorage: %v", err)
|
||||
}
|
||||
defer storage.Close()
|
||||
for _, dir := range []string{sessionsDir, streamsDir} {
|
||||
if err := storage.Delete(context.Background(), dir); err != nil {
|
||||
t.Fatalf("clearing %s: %v", dir, err)
|
||||
}
|
||||
}
|
||||
|
||||
return config
|
||||
}
|
||||
|
||||
func envInt(t *testing.T, name string) int {
|
||||
t.Helper()
|
||||
|
||||
parsed, err := strconv.Atoi(os.Getenv(name))
|
||||
if err != nil {
|
||||
t.Fatalf("%s: %v", name, err)
|
||||
}
|
||||
return parsed
|
||||
}
|
||||
|
||||
func TestLiveDriveStorage(t *testing.T) {
|
||||
config := liveDriveConfig(t)
|
||||
storage, err := newDriveStorage(nil, config)
|
||||
if err != nil {
|
||||
t.Fatalf("newDriveStorage: %v", err)
|
||||
}
|
||||
defer storage.Close()
|
||||
|
||||
ctx := context.Background()
|
||||
name := "streams/livetest/c2s/000000000.seg"
|
||||
payload := []byte("xdrive over a real remote storage service")
|
||||
defer storage.Delete(ctx, "streams/livetest")
|
||||
|
||||
start := time.Now()
|
||||
if err := storage.Put(ctx, name, payload); err != nil {
|
||||
t.Fatalf("Put: %v", err)
|
||||
}
|
||||
t.Logf("Put took %v", time.Since(start))
|
||||
|
||||
start = time.Now()
|
||||
names, err := storage.List(ctx, "streams/livetest/c2s")
|
||||
if err != nil {
|
||||
t.Fatalf("List: %v", err)
|
||||
}
|
||||
t.Logf("List took %v", time.Since(start))
|
||||
if len(names) != 1 || names[0].Name != "000000000.seg" {
|
||||
t.Fatalf("List returned %v, want one segment", names)
|
||||
}
|
||||
|
||||
start = time.Now()
|
||||
got, err := storage.Get(ctx, name)
|
||||
if err != nil {
|
||||
t.Fatalf("Get: %v", err)
|
||||
}
|
||||
t.Logf("Get took %v", time.Since(start))
|
||||
if !bytes.Equal(got, payload) {
|
||||
t.Fatalf("Get returned %q, want %q", got, payload)
|
||||
}
|
||||
|
||||
if _, err := storage.Get(ctx, "streams/livetest/c2s/000000009.seg"); err != errNotFound {
|
||||
t.Fatalf("Get returned %v, want errNotFound", err)
|
||||
}
|
||||
|
||||
if err := storage.Delete(ctx, "streams/livetest"); err != nil {
|
||||
t.Fatalf("Delete: %v", err)
|
||||
}
|
||||
names, err = storage.List(ctx, "streams/livetest/c2s")
|
||||
if err != nil {
|
||||
t.Fatalf("List after delete: %v", err)
|
||||
}
|
||||
if len(names) != 0 {
|
||||
t.Fatalf("List after delete returned %v", names)
|
||||
}
|
||||
}
|
||||
|
||||
func TestLiveDriveTransport(t *testing.T) {
|
||||
config := liveDriveConfig(t)
|
||||
streamSettings := &internet.MemoryStreamConfig{
|
||||
ProtocolName: protocolName,
|
||||
ProtocolSettings: config,
|
||||
}
|
||||
|
||||
client, server, cleanup := pairWith(t, streamSettings)
|
||||
defer cleanup()
|
||||
|
||||
start := time.Now()
|
||||
if _, err := client.Write([]byte("ping")); err != nil {
|
||||
t.Fatalf("client write: %v", err)
|
||||
}
|
||||
expectRead(t, server, "ping")
|
||||
t.Logf("client to server round took %v", time.Since(start))
|
||||
|
||||
start = time.Now()
|
||||
if _, err := server.Write([]byte("pong")); err != nil {
|
||||
t.Fatalf("server write: %v", err)
|
||||
}
|
||||
expectRead(t, client, "pong")
|
||||
t.Logf("server to client round took %v", time.Since(start))
|
||||
|
||||
size := 400000
|
||||
if raw := os.Getenv("XRAY_XDRIVE_LIVE_BYTES"); raw != "" {
|
||||
parsed, err := strconv.Atoi(raw)
|
||||
if err != nil {
|
||||
t.Fatalf("XRAY_XDRIVE_LIVE_BYTES: %v", err)
|
||||
}
|
||||
size = parsed
|
||||
}
|
||||
payload := make([]byte, size)
|
||||
if _, err := rand.Read(payload); err != nil {
|
||||
t.Fatalf("rand: %v", err)
|
||||
}
|
||||
|
||||
start = time.Now()
|
||||
go func() {
|
||||
client.Write(payload)
|
||||
}()
|
||||
if err := server.SetReadDeadline(time.Now().Add(5 * time.Minute)); err != nil {
|
||||
t.Fatalf("SetReadDeadline: %v", err)
|
||||
}
|
||||
got := make([]byte, len(payload))
|
||||
if _, err := io.ReadFull(server, got); err != nil {
|
||||
t.Fatalf("ReadFull: %v", err)
|
||||
}
|
||||
elapsed := time.Since(start)
|
||||
if !bytes.Equal(got, payload) {
|
||||
t.Fatal("payload mismatch")
|
||||
}
|
||||
t.Logf("%d bytes took %v (%.1f KiB/s)", len(payload), elapsed,
|
||||
float64(len(payload))/1024/elapsed.Seconds())
|
||||
}
|
||||
|
||||
func TestLiveDriveParallelPut(t *testing.T) {
|
||||
config := liveDriveConfig(t)
|
||||
storage, err := newDriveStorage(nil, config)
|
||||
if err != nil {
|
||||
t.Fatalf("newDriveStorage: %v", err)
|
||||
}
|
||||
defer storage.Close()
|
||||
|
||||
ctx := context.Background()
|
||||
defer storage.Delete(ctx, "streams/benchtest")
|
||||
|
||||
chunk := make([]byte, 256*1024)
|
||||
if _, err := rand.Read(chunk); err != nil {
|
||||
t.Fatalf("rand: %v", err)
|
||||
}
|
||||
|
||||
if err := storage.Put(ctx, "streams/benchtest/warmup", chunk); err != nil {
|
||||
t.Fatalf("warmup: %v", err)
|
||||
}
|
||||
|
||||
start := time.Now()
|
||||
for i := 0; i < 4; i++ {
|
||||
if err := storage.Put(ctx, fmt.Sprintf("streams/benchtest/seq%d", i), chunk); err != nil {
|
||||
t.Fatalf("sequential put: %v", err)
|
||||
}
|
||||
}
|
||||
sequential := time.Since(start)
|
||||
t.Logf("4 sequential puts of 256 KiB: %v (%.1f KiB/s)",
|
||||
sequential, float64(4*len(chunk))/1024/sequential.Seconds())
|
||||
|
||||
start = time.Now()
|
||||
var wg sync.WaitGroup
|
||||
failures := make([]error, 8)
|
||||
for i := 0; i < 8; i++ {
|
||||
wg.Add(1)
|
||||
go func(i int) {
|
||||
defer wg.Done()
|
||||
failures[i] = storage.Put(ctx, fmt.Sprintf("streams/benchtest/par%d", i), chunk)
|
||||
}(i)
|
||||
}
|
||||
wg.Wait()
|
||||
parallel := time.Since(start)
|
||||
for _, err := range failures {
|
||||
if err != nil {
|
||||
t.Fatalf("parallel put: %v", err)
|
||||
}
|
||||
}
|
||||
t.Logf("8 parallel puts of 256 KiB: %v (%.1f KiB/s)",
|
||||
parallel, float64(8*len(chunk))/1024/parallel.Seconds())
|
||||
|
||||
start = time.Now()
|
||||
names, err := storage.List(ctx, "streams/benchtest")
|
||||
if err != nil {
|
||||
t.Fatalf("List: %v", err)
|
||||
}
|
||||
t.Logf("List of %d objects took %v", len(names), time.Since(start))
|
||||
|
||||
start = time.Now()
|
||||
wg = sync.WaitGroup{}
|
||||
for i := 0; i < 8; i++ {
|
||||
wg.Add(1)
|
||||
go func(i int) {
|
||||
defer wg.Done()
|
||||
storage.Get(ctx, fmt.Sprintf("streams/benchtest/par%d", i))
|
||||
}(i)
|
||||
}
|
||||
wg.Wait()
|
||||
download := time.Since(start)
|
||||
t.Logf("8 parallel gets of 256 KiB: %v (%.1f KiB/s)",
|
||||
download, float64(8*len(chunk))/1024/download.Seconds())
|
||||
}
|
||||
|
||||
func TestLiveDriveSegmentSweep(t *testing.T) {
|
||||
config := liveDriveConfig(t)
|
||||
storage, err := newDriveStorage(nil, config)
|
||||
if err != nil {
|
||||
t.Fatalf("newDriveStorage: %v", err)
|
||||
}
|
||||
defer storage.Close()
|
||||
|
||||
ctx := context.Background()
|
||||
defer storage.Delete(ctx, "streams/sweeptest")
|
||||
|
||||
const total = 1024 * 1024
|
||||
for _, size := range []int{64 * 1024, 128 * 1024, 256 * 1024, 512 * 1024} {
|
||||
chunk := make([]byte, size)
|
||||
if _, err := rand.Read(chunk); err != nil {
|
||||
t.Fatalf("rand: %v", err)
|
||||
}
|
||||
count := total / size
|
||||
|
||||
start := time.Now()
|
||||
var wg sync.WaitGroup
|
||||
for i := 0; i < count; i++ {
|
||||
wg.Add(1)
|
||||
go func(i int) {
|
||||
defer wg.Done()
|
||||
storage.Put(ctx, fmt.Sprintf("streams/sweeptest/s%d-%d", size, i), chunk)
|
||||
}(i)
|
||||
}
|
||||
wg.Wait()
|
||||
elapsed := time.Since(start)
|
||||
|
||||
t.Logf("%4d KiB x %2d = 1 MiB in %8v -> %6.1f KiB/s",
|
||||
size/1024, count, elapsed.Round(time.Millisecond),
|
||||
float64(total)/1024/elapsed.Seconds())
|
||||
}
|
||||
}
|
||||
|
||||
func TestLiveDriveListLag(t *testing.T) {
|
||||
config := liveDriveConfig(t)
|
||||
storage, err := newDriveStorage(nil, config)
|
||||
if err != nil {
|
||||
t.Fatalf("newDriveStorage: %v", err)
|
||||
}
|
||||
defer storage.Close()
|
||||
|
||||
ctx := context.Background()
|
||||
defer storage.Delete(ctx, "streams/lagtest")
|
||||
|
||||
const rounds = 6
|
||||
var worst time.Duration
|
||||
|
||||
for i := 0; i < rounds; i++ {
|
||||
name := fmt.Sprintf("streams/lagtest/round%d/000000000.seg", i)
|
||||
if err := storage.Put(ctx, name, []byte("probe")); err != nil {
|
||||
t.Fatalf("Put: %v", err)
|
||||
}
|
||||
|
||||
start := time.Now()
|
||||
var lag time.Duration
|
||||
for {
|
||||
names, err := storage.List(ctx, fmt.Sprintf("streams/lagtest/round%d", i))
|
||||
if err != nil {
|
||||
t.Fatalf("List: %v", err)
|
||||
}
|
||||
if len(names) == 1 {
|
||||
lag = time.Since(start)
|
||||
break
|
||||
}
|
||||
if time.Since(start) > 30*time.Second {
|
||||
t.Fatalf("round %d: the object never showed up in a listing", i)
|
||||
}
|
||||
}
|
||||
if lag > worst {
|
||||
worst = lag
|
||||
}
|
||||
t.Logf("round %d: the object became listable after %v", i, lag.Round(time.Millisecond))
|
||||
}
|
||||
t.Logf("worst listing lag: %v", worst.Round(time.Millisecond))
|
||||
}
|
||||
@@ -0,0 +1,690 @@
|
||||
package xdrive
|
||||
|
||||
import (
|
||||
"context"
|
||||
"encoding/json"
|
||||
"fmt"
|
||||
"io"
|
||||
"mime"
|
||||
"mime/multipart"
|
||||
"net/http"
|
||||
"net/http/httptest"
|
||||
"net/url"
|
||||
"strings"
|
||||
"sync"
|
||||
"testing"
|
||||
"time"
|
||||
|
||||
"github.com/xtls/xray-core/common/net"
|
||||
"github.com/xtls/xray-core/transport/internet"
|
||||
)
|
||||
|
||||
type fakeFile struct {
|
||||
id string
|
||||
name string
|
||||
data []byte
|
||||
description string
|
||||
}
|
||||
|
||||
type fakeDrive struct {
|
||||
server *httptest.Server
|
||||
|
||||
mu sync.Mutex
|
||||
files map[string]*fakeFile
|
||||
nextID int
|
||||
failOnce map[string]bool
|
||||
failStatus map[string]int
|
||||
failBody map[string]string
|
||||
tokens int
|
||||
hosts map[string]bool
|
||||
}
|
||||
|
||||
func newFakeDrive(t *testing.T) *fakeDrive {
|
||||
t.Helper()
|
||||
|
||||
drive := &fakeDrive{
|
||||
files: make(map[string]*fakeFile),
|
||||
failOnce: make(map[string]bool),
|
||||
failStatus: make(map[string]int),
|
||||
failBody: make(map[string]string),
|
||||
hosts: make(map[string]bool),
|
||||
}
|
||||
mux := http.NewServeMux()
|
||||
mux.Handle("/", http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) {
|
||||
http.NotFound(w, r)
|
||||
}))
|
||||
record := func(next http.HandlerFunc) http.HandlerFunc {
|
||||
return func(w http.ResponseWriter, r *http.Request) {
|
||||
drive.mu.Lock()
|
||||
drive.hosts[r.Host] = true
|
||||
drive.mu.Unlock()
|
||||
next(w, r)
|
||||
}
|
||||
}
|
||||
mux.HandleFunc("/token", record(drive.handleToken))
|
||||
mux.HandleFunc("/upload", record(drive.handleUpload))
|
||||
mux.HandleFunc("/files", record(drive.handleFiles))
|
||||
mux.HandleFunc("/files/", record(drive.handleFile))
|
||||
drive.server = httptest.NewServer(mux)
|
||||
|
||||
resetSharedStorage()
|
||||
|
||||
previous := []string{driveTokenURL, driveFilesURL, driveUploadURL}
|
||||
previousBackoff := driveInitialBackoff
|
||||
driveTokenURL = drive.server.URL + "/token"
|
||||
driveFilesURL = drive.server.URL + "/files"
|
||||
driveUploadURL = drive.server.URL + "/upload"
|
||||
driveInitialBackoff = 5 * time.Millisecond
|
||||
|
||||
t.Cleanup(func() {
|
||||
driveTokenURL, driveFilesURL, driveUploadURL = previous[0], previous[1], previous[2]
|
||||
driveInitialBackoff = previousBackoff
|
||||
drive.server.Close()
|
||||
resetSharedStorage()
|
||||
})
|
||||
return drive
|
||||
}
|
||||
|
||||
func (d *fakeDrive) handleToken(w http.ResponseWriter, r *http.Request) {
|
||||
d.mu.Lock()
|
||||
d.tokens++
|
||||
d.mu.Unlock()
|
||||
json.NewEncoder(w).Encode(map[string]interface{}{
|
||||
"access_token": "fake-token",
|
||||
"expires_in": 3600,
|
||||
})
|
||||
}
|
||||
|
||||
func (d *fakeDrive) handleUpload(w http.ResponseWriter, r *http.Request) {
|
||||
if d.shouldFail("upload") {
|
||||
if status, body := d.failure("upload"); status != 0 {
|
||||
w.WriteHeader(status)
|
||||
w.Write([]byte(body))
|
||||
return
|
||||
}
|
||||
w.WriteHeader(http.StatusTooManyRequests)
|
||||
return
|
||||
}
|
||||
|
||||
_, params, err := mime.ParseMediaType(r.Header.Get("Content-Type"))
|
||||
if err != nil {
|
||||
w.WriteHeader(http.StatusBadRequest)
|
||||
return
|
||||
}
|
||||
reader := multipart.NewReader(r.Body, params["boundary"])
|
||||
|
||||
metaPart, err := reader.NextPart()
|
||||
if err != nil {
|
||||
w.WriteHeader(http.StatusBadRequest)
|
||||
return
|
||||
}
|
||||
var metadata struct {
|
||||
Name string `json:"name"`
|
||||
}
|
||||
if err := json.NewDecoder(metaPart).Decode(&metadata); err != nil {
|
||||
w.WriteHeader(http.StatusBadRequest)
|
||||
return
|
||||
}
|
||||
|
||||
var data []byte
|
||||
if dataPart, err := reader.NextPart(); err == nil {
|
||||
data, _ = io.ReadAll(dataPart)
|
||||
}
|
||||
|
||||
d.mu.Lock()
|
||||
d.nextID++
|
||||
id := fmt.Sprintf("id-%d", d.nextID)
|
||||
d.files[id] = &fakeFile{id: id, name: metadata.Name, data: data}
|
||||
d.mu.Unlock()
|
||||
|
||||
json.NewEncoder(w).Encode(map[string]string{"id": id, "name": metadata.Name})
|
||||
}
|
||||
|
||||
func (d *fakeDrive) handleFiles(w http.ResponseWriter, r *http.Request) {
|
||||
if r.Method == http.MethodPost {
|
||||
d.handleCreate(w, r)
|
||||
return
|
||||
}
|
||||
d.handleList(w, r)
|
||||
}
|
||||
|
||||
func (d *fakeDrive) handleCreate(w http.ResponseWriter, r *http.Request) {
|
||||
if d.shouldFail("upload") {
|
||||
if status, body := d.failure("upload"); status != 0 {
|
||||
w.WriteHeader(status)
|
||||
w.Write([]byte(body))
|
||||
return
|
||||
}
|
||||
w.WriteHeader(http.StatusTooManyRequests)
|
||||
return
|
||||
}
|
||||
|
||||
var meta struct {
|
||||
Name string `json:"name"`
|
||||
Description string `json:"description"`
|
||||
}
|
||||
if err := json.NewDecoder(r.Body).Decode(&meta); err != nil {
|
||||
w.WriteHeader(http.StatusBadRequest)
|
||||
return
|
||||
}
|
||||
|
||||
d.mu.Lock()
|
||||
d.nextID++
|
||||
id := fmt.Sprintf("id-%d", d.nextID)
|
||||
d.files[id] = &fakeFile{id: id, name: meta.Name, description: meta.Description}
|
||||
d.mu.Unlock()
|
||||
|
||||
json.NewEncoder(w).Encode(map[string]string{"id": id, "name": meta.Name})
|
||||
}
|
||||
|
||||
func (d *fakeDrive) handleList(w http.ResponseWriter, r *http.Request) {
|
||||
if d.shouldFail("list") {
|
||||
w.WriteHeader(http.StatusServiceUnavailable)
|
||||
return
|
||||
}
|
||||
|
||||
query := r.URL.Query().Get("q")
|
||||
exact, prefix := parseFakeQuery(query)
|
||||
|
||||
type entry struct {
|
||||
ID string `json:"id"`
|
||||
Name string `json:"name"`
|
||||
Description string `json:"description,omitempty"`
|
||||
}
|
||||
result := struct {
|
||||
Files []entry `json:"files"`
|
||||
}{}
|
||||
|
||||
d.mu.Lock()
|
||||
for _, file := range d.files {
|
||||
match := false
|
||||
switch {
|
||||
case exact != "":
|
||||
match = file.name == exact
|
||||
case prefix != "":
|
||||
match = strings.HasPrefix(file.name, prefix)
|
||||
}
|
||||
if match {
|
||||
result.Files = append(result.Files, entry{
|
||||
ID: file.id, Name: file.name, Description: file.description,
|
||||
})
|
||||
}
|
||||
}
|
||||
d.mu.Unlock()
|
||||
|
||||
json.NewEncoder(w).Encode(result)
|
||||
}
|
||||
|
||||
func (d *fakeDrive) handleFile(w http.ResponseWriter, r *http.Request) {
|
||||
id := strings.TrimPrefix(r.URL.Path, "/files/")
|
||||
|
||||
d.mu.Lock()
|
||||
file, ok := d.files[id]
|
||||
if ok && r.Method == http.MethodDelete {
|
||||
delete(d.files, id)
|
||||
}
|
||||
d.mu.Unlock()
|
||||
|
||||
if !ok {
|
||||
w.WriteHeader(http.StatusNotFound)
|
||||
return
|
||||
}
|
||||
if r.Method == http.MethodDelete {
|
||||
w.WriteHeader(http.StatusNoContent)
|
||||
return
|
||||
}
|
||||
if strings.Contains(r.URL.RawQuery, "fields=description") {
|
||||
json.NewEncoder(w).Encode(map[string]string{"description": file.description})
|
||||
return
|
||||
}
|
||||
w.Write(file.data)
|
||||
}
|
||||
|
||||
func (d *fakeDrive) shouldFail(kind string) bool {
|
||||
d.mu.Lock()
|
||||
defer d.mu.Unlock()
|
||||
if d.failOnce[kind] {
|
||||
d.failOnce[kind] = false
|
||||
return true
|
||||
}
|
||||
return false
|
||||
}
|
||||
|
||||
func (d *fakeDrive) failOnceWith(kind string, status int, body string) {
|
||||
d.mu.Lock()
|
||||
d.failOnce[kind] = true
|
||||
d.failStatus[kind] = status
|
||||
d.failBody[kind] = body
|
||||
d.mu.Unlock()
|
||||
}
|
||||
|
||||
func (d *fakeDrive) failure(kind string) (int, string) {
|
||||
d.mu.Lock()
|
||||
defer d.mu.Unlock()
|
||||
if status, ok := d.failStatus[kind]; ok {
|
||||
return status, d.failBody[kind]
|
||||
}
|
||||
return 0, ""
|
||||
}
|
||||
|
||||
func (d *fakeDrive) failNext(kind string) {
|
||||
d.mu.Lock()
|
||||
d.failOnce[kind] = true
|
||||
d.mu.Unlock()
|
||||
}
|
||||
|
||||
func (d *fakeDrive) seenHosts() []string {
|
||||
d.mu.Lock()
|
||||
defer d.mu.Unlock()
|
||||
out := make([]string, 0, len(d.hosts))
|
||||
for h := range d.hosts {
|
||||
out = append(out, h)
|
||||
}
|
||||
return out
|
||||
}
|
||||
|
||||
func (d *fakeDrive) count() int {
|
||||
d.mu.Lock()
|
||||
defer d.mu.Unlock()
|
||||
return len(d.files)
|
||||
}
|
||||
|
||||
func parseFakeQuery(query string) (exact, prefix string) {
|
||||
if value, ok := cutQuoted(query, "name = '"); ok {
|
||||
return value, ""
|
||||
}
|
||||
if value, ok := cutQuoted(query, "name contains '"); ok {
|
||||
return "", value
|
||||
}
|
||||
return "", ""
|
||||
}
|
||||
|
||||
func cutQuoted(query, marker string) (string, bool) {
|
||||
start := strings.Index(query, marker)
|
||||
if start < 0 {
|
||||
return "", false
|
||||
}
|
||||
rest := query[start+len(marker):]
|
||||
end := strings.Index(rest, "'")
|
||||
if end < 0 {
|
||||
return "", false
|
||||
}
|
||||
return rest[:end], true
|
||||
}
|
||||
|
||||
func driveSettings() *internet.MemoryStreamConfig {
|
||||
return &internet.MemoryStreamConfig{
|
||||
ProtocolName: protocolName,
|
||||
ProtocolSettings: &Config{
|
||||
RemoteFolder: "folder-id",
|
||||
Service: "Google Drive",
|
||||
Secrets: []string{"client", "secret", "refresh"},
|
||||
FlushIntervalMs: 5,
|
||||
PollIntervalMs: 5,
|
||||
MaxPollIntervalMs: 20,
|
||||
SessionTtlSeconds: 1,
|
||||
},
|
||||
}
|
||||
}
|
||||
|
||||
func newDriveBackend(t *testing.T) *driveStorage {
|
||||
t.Helper()
|
||||
|
||||
storage, err := newDriveStorage(driveSettings(), driveSettings().ProtocolSettings.(*Config))
|
||||
if err != nil {
|
||||
t.Fatalf("newDriveStorage: %v", err)
|
||||
}
|
||||
return storage
|
||||
}
|
||||
|
||||
func TestDriveSecrets(t *testing.T) {
|
||||
if _, err := newDriveStorage(nil, &Config{RemoteFolder: "f", Secrets: []string{"a", "b"}}); err == nil {
|
||||
t.Fatal("accepted two secrets")
|
||||
}
|
||||
if _, err := newDriveStorage(nil, &Config{Secrets: []string{"a", "b", "c"}}); err == nil {
|
||||
t.Fatal("accepted an empty remoteFolder")
|
||||
}
|
||||
if _, err := newDriveStorage(nil, &Config{RemoteFolder: "f", Secrets: []string{"a", "", "c"}}); err == nil {
|
||||
t.Fatal("accepted an empty secret")
|
||||
}
|
||||
}
|
||||
|
||||
func TestDriveRoundTrip(t *testing.T) {
|
||||
drive := newFakeDrive(t)
|
||||
storage := newDriveBackend(t)
|
||||
ctx := context.Background()
|
||||
|
||||
if err := storage.Put(ctx, "streams/abc/c2s/000000000.seg", []byte("hello")); err != nil {
|
||||
t.Fatalf("Put: %v", err)
|
||||
}
|
||||
|
||||
data, err := storage.Get(ctx, "streams/abc/c2s/000000000.seg")
|
||||
if err != nil {
|
||||
t.Fatalf("Get: %v", err)
|
||||
}
|
||||
if string(data) != "hello" {
|
||||
t.Fatalf("Get returned %q, want %q", data, "hello")
|
||||
}
|
||||
|
||||
names, err := storage.List(ctx, "streams/abc/c2s")
|
||||
if err != nil {
|
||||
t.Fatalf("List: %v", err)
|
||||
}
|
||||
if len(names) != 1 || names[0].Name != "000000000.seg" {
|
||||
t.Fatalf("List returned %v, want 1 segment", names)
|
||||
}
|
||||
|
||||
if _, err := storage.Get(ctx, "streams/abc/c2s/000000009.seg"); err != errNotFound {
|
||||
t.Fatalf("Get returned %v, want errNotFound", err)
|
||||
}
|
||||
|
||||
if err := storage.Delete(ctx, "streams/abc/c2s/000000000.seg"); err != nil {
|
||||
t.Fatalf("Delete: %v", err)
|
||||
}
|
||||
if drive.count() != 0 {
|
||||
t.Fatalf("fake drive still holds %d files", drive.count())
|
||||
}
|
||||
}
|
||||
|
||||
func TestDriveListChildren(t *testing.T) {
|
||||
newFakeDrive(t)
|
||||
storage := newDriveBackend(t)
|
||||
ctx := context.Background()
|
||||
|
||||
for _, name := range []string{
|
||||
"streams/one/c2s/000000000.seg",
|
||||
"streams/one/s2c/000000000.seg",
|
||||
"streams/two/c2s/000000000.seg",
|
||||
} {
|
||||
if err := storage.Put(ctx, name, []byte("x")); err != nil {
|
||||
t.Fatalf("Put %s: %v", name, err)
|
||||
}
|
||||
}
|
||||
|
||||
names, err := storage.List(ctx, "streams")
|
||||
if err != nil {
|
||||
t.Fatalf("List: %v", err)
|
||||
}
|
||||
if len(names) != 2 {
|
||||
t.Fatalf("List returned %v, want 2 sessions", names)
|
||||
}
|
||||
}
|
||||
|
||||
func TestDriveDeleteSession(t *testing.T) {
|
||||
drive := newFakeDrive(t)
|
||||
storage := newDriveBackend(t)
|
||||
ctx := context.Background()
|
||||
|
||||
for _, name := range []string{
|
||||
"streams/one/c2s/000000000.seg",
|
||||
"streams/one/c2s/000000001.end",
|
||||
"streams/one/s2c/000000000.seg",
|
||||
"streams/two/c2s/000000000.seg",
|
||||
} {
|
||||
if err := storage.Put(ctx, name, []byte("x")); err != nil {
|
||||
t.Fatalf("Put %s: %v", name, err)
|
||||
}
|
||||
}
|
||||
|
||||
if err := storage.Delete(ctx, "streams/one"); err != nil {
|
||||
t.Fatalf("Delete: %v", err)
|
||||
}
|
||||
if drive.count() != 1 {
|
||||
t.Fatalf("fake drive holds %d files, want 1", drive.count())
|
||||
}
|
||||
}
|
||||
|
||||
func TestDriveRetry(t *testing.T) {
|
||||
drive := newFakeDrive(t)
|
||||
storage := newDriveBackend(t)
|
||||
ctx := context.Background()
|
||||
|
||||
drive.failNext("upload")
|
||||
if err := storage.Put(ctx, "sessions/abc", nil); err != nil {
|
||||
t.Fatalf("Put did not survive a 429: %v", err)
|
||||
}
|
||||
|
||||
drive.failNext("list")
|
||||
if _, err := storage.List(ctx, "sessions"); err != nil {
|
||||
t.Fatalf("List did not survive a 503: %v", err)
|
||||
}
|
||||
}
|
||||
|
||||
func TestDriveTokenCache(t *testing.T) {
|
||||
drive := newFakeDrive(t)
|
||||
storage := newDriveBackend(t)
|
||||
ctx := context.Background()
|
||||
|
||||
for i := 0; i < 5; i++ {
|
||||
if err := storage.Put(ctx, fmt.Sprintf("sessions/s%d", i), nil); err != nil {
|
||||
t.Fatalf("Put: %v", err)
|
||||
}
|
||||
}
|
||||
if drive.tokens != 1 {
|
||||
t.Fatalf("token endpoint hit %d times, want 1", drive.tokens)
|
||||
}
|
||||
}
|
||||
|
||||
func TestDriveTransport(t *testing.T) {
|
||||
newFakeDrive(t)
|
||||
|
||||
client, server, cleanup := pairWith(t, driveSettings())
|
||||
defer cleanup()
|
||||
|
||||
if _, err := client.Write([]byte("ping")); err != nil {
|
||||
t.Fatalf("client write: %v", err)
|
||||
}
|
||||
expectRead(t, server, "ping")
|
||||
|
||||
if _, err := server.Write([]byte("pong")); err != nil {
|
||||
t.Fatalf("server write: %v", err)
|
||||
}
|
||||
expectRead(t, client, "pong")
|
||||
}
|
||||
|
||||
func TestDriveLargeTransfer(t *testing.T) {
|
||||
newFakeDrive(t)
|
||||
|
||||
client, server, cleanup := pairWith(t, driveSettings())
|
||||
defer cleanup()
|
||||
|
||||
payload := make([]byte, 300000)
|
||||
for i := range payload {
|
||||
payload[i] = byte(i % 251)
|
||||
}
|
||||
|
||||
go func() {
|
||||
client.Write(payload)
|
||||
}()
|
||||
|
||||
if err := server.SetReadDeadline(time.Now().Add(60 * time.Second)); err != nil {
|
||||
t.Fatalf("SetReadDeadline: %v", err)
|
||||
}
|
||||
got := make([]byte, len(payload))
|
||||
if _, err := io.ReadFull(server, got); err != nil {
|
||||
t.Fatalf("ReadFull: %v", err)
|
||||
}
|
||||
for i := range got {
|
||||
if got[i] != payload[i] {
|
||||
t.Fatalf("payload mismatch at byte %d", i)
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
func TestRateLimited(t *testing.T) {
|
||||
limited := []string{
|
||||
`{"error":{"code":403,"errors":[{"reason":"userRateLimitExceeded"}]}}`,
|
||||
`{"error":{"code":403,"errors":[{"reason":"rateLimitExceeded"}]}}`,
|
||||
`{"error":{"code":403,"errors":[{"reason":"sharingRateLimitExceeded"}]}}`,
|
||||
`{"error":{"status":"RESOURCE_EXHAUSTED"}}`,
|
||||
}
|
||||
for _, payload := range limited {
|
||||
if !rateLimited([]byte(payload)) {
|
||||
t.Fatalf("rateLimited missed %s", payload)
|
||||
}
|
||||
}
|
||||
|
||||
permanent := []string{
|
||||
`{"error":{"code":403,"errors":[{"reason":"insufficientFilePermissions"}]}}`,
|
||||
`{"error":{"code":403,"errors":[{"reason":"storageQuotaExceeded"}]}}`,
|
||||
`not json at all`,
|
||||
}
|
||||
for _, payload := range permanent {
|
||||
if rateLimited([]byte(payload)) {
|
||||
t.Fatalf("rateLimited treated %s as temporary", payload)
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
func TestDriveRetryRateLimit(t *testing.T) {
|
||||
drive := newFakeDrive(t)
|
||||
storage := newDriveBackend(t)
|
||||
|
||||
drive.failOnceWith("upload", http.StatusForbidden,
|
||||
`{"error":{"code":403,"errors":[{"reason":"userRateLimitExceeded"}]}}`)
|
||||
|
||||
if err := storage.Put(context.Background(), "sessions/abc", nil); err != nil {
|
||||
t.Fatalf("Put did not survive a 403: %v", err)
|
||||
}
|
||||
}
|
||||
|
||||
func TestDriveSharedClient(t *testing.T) {
|
||||
newFakeDrive(t)
|
||||
|
||||
first, err := newStorage(driveSettings())
|
||||
if err != nil {
|
||||
t.Fatalf("newStorage: %v", err)
|
||||
}
|
||||
second, err := newStorage(driveSettings())
|
||||
if err != nil {
|
||||
t.Fatalf("newStorage: %v", err)
|
||||
}
|
||||
if first != second {
|
||||
t.Fatal("same settings did not share one storage")
|
||||
}
|
||||
}
|
||||
|
||||
func TestDriveInlineListing(t *testing.T) {
|
||||
drive := newFakeDrive(t)
|
||||
storage := newDriveBackend(t)
|
||||
ctx := context.Background()
|
||||
|
||||
payload := []byte("small enough to ride along with the listing")
|
||||
if err := storage.Put(ctx, "streams/abc/c2s/000000000.seg", payload); err != nil {
|
||||
t.Fatalf("Put: %v", err)
|
||||
}
|
||||
|
||||
entries, err := storage.List(ctx, "streams/abc/c2s")
|
||||
if err != nil {
|
||||
t.Fatalf("List: %v", err)
|
||||
}
|
||||
if len(entries) != 1 {
|
||||
t.Fatalf("List returned %v, want 1 entry", entries)
|
||||
}
|
||||
if string(entries[0].Inline) != string(payload) {
|
||||
t.Fatalf("listing carried %q, want %q", entries[0].Inline, payload)
|
||||
}
|
||||
|
||||
d := drive
|
||||
d.mu.Lock()
|
||||
var stored *fakeFile
|
||||
for _, f := range d.files {
|
||||
stored = f
|
||||
}
|
||||
d.mu.Unlock()
|
||||
if len(stored.data) != 0 {
|
||||
t.Fatal("small payload was uploaded as content")
|
||||
}
|
||||
}
|
||||
|
||||
func TestDriveLargeContent(t *testing.T) {
|
||||
drive := newFakeDrive(t)
|
||||
storage := newDriveBackend(t)
|
||||
ctx := context.Background()
|
||||
|
||||
payload := make([]byte, driveInlineLimit+1)
|
||||
for i := range payload {
|
||||
payload[i] = byte(i)
|
||||
}
|
||||
if err := storage.Put(ctx, "streams/abc/c2s/000000000.seg", payload); err != nil {
|
||||
t.Fatalf("Put: %v", err)
|
||||
}
|
||||
|
||||
entries, err := storage.List(ctx, "streams/abc/c2s")
|
||||
if err != nil {
|
||||
t.Fatalf("List: %v", err)
|
||||
}
|
||||
if len(entries) != 1 || entries[0].Inline != nil {
|
||||
t.Fatalf("large payload was inlined, got %v", entries)
|
||||
}
|
||||
|
||||
got, err := storage.Get(ctx, "streams/abc/c2s/000000000.seg")
|
||||
if err != nil {
|
||||
t.Fatalf("Get: %v", err)
|
||||
}
|
||||
if len(got) != len(payload) {
|
||||
t.Fatalf("Get returned %d bytes, want %d", len(got), len(payload))
|
||||
}
|
||||
_ = drive
|
||||
}
|
||||
|
||||
func TestDriveInlineGet(t *testing.T) {
|
||||
newFakeDrive(t)
|
||||
storage := newDriveBackend(t)
|
||||
ctx := context.Background()
|
||||
|
||||
payload := []byte("only in the description")
|
||||
if err := storage.Put(ctx, "sessions/abc", payload); err != nil {
|
||||
t.Fatalf("Put: %v", err)
|
||||
}
|
||||
|
||||
got, err := storage.Get(ctx, "sessions/abc")
|
||||
if err != nil {
|
||||
t.Fatalf("Get: %v", err)
|
||||
}
|
||||
if string(got) != string(payload) {
|
||||
t.Fatalf("Get returned %q, want %q", got, payload)
|
||||
}
|
||||
}
|
||||
|
||||
func TestDriveFronting(t *testing.T) {
|
||||
drive := newFakeDrive(t)
|
||||
|
||||
fake, err := url.Parse(drive.server.URL)
|
||||
if err != nil {
|
||||
t.Fatalf("parse: %v", err)
|
||||
}
|
||||
port, err := net.PortFromString(fake.Port())
|
||||
if err != nil {
|
||||
t.Fatalf("port: %v", err)
|
||||
}
|
||||
|
||||
driveTokenURL = "http://www.googleapis.com/token"
|
||||
driveFilesURL = "http://www.googleapis.com/files"
|
||||
driveUploadURL = "http://www.googleapis.com/upload"
|
||||
|
||||
settings := driveSettings()
|
||||
settings.Destination = &net.Destination{
|
||||
Address: net.ParseAddress(fake.Hostname()),
|
||||
Port: port,
|
||||
Network: net.Network_TCP,
|
||||
}
|
||||
|
||||
storage, err := newDriveStorage(settings, settings.ProtocolSettings.(*Config))
|
||||
if err != nil {
|
||||
t.Fatalf("newDriveStorage: %v", err)
|
||||
}
|
||||
|
||||
if err := storage.Put(context.Background(), "sessions/fronted", []byte("x")); err != nil {
|
||||
t.Fatalf("Put: %v", err)
|
||||
}
|
||||
if drive.count() != 1 {
|
||||
t.Fatalf("fake drive holds %d files, want 1", drive.count())
|
||||
}
|
||||
|
||||
for _, host := range drive.seenHosts() {
|
||||
if host != "www.googleapis.com" {
|
||||
t.Fatalf("inner host was %q, want www.googleapis.com regardless of address", host)
|
||||
}
|
||||
}
|
||||
}
|
||||
@@ -0,0 +1,39 @@
|
||||
package xdrive
|
||||
|
||||
import (
|
||||
"encoding/json"
|
||||
"strings"
|
||||
)
|
||||
|
||||
func jsonWalk(payload []byte, path string) interface{} {
|
||||
var root interface{}
|
||||
if json.Unmarshal(payload, &root) != nil {
|
||||
return nil
|
||||
}
|
||||
node := root
|
||||
for _, key := range strings.Split(path, ".") {
|
||||
obj, ok := node.(map[string]interface{})
|
||||
if !ok {
|
||||
return nil
|
||||
}
|
||||
node, ok = obj[key]
|
||||
if !ok {
|
||||
return nil
|
||||
}
|
||||
}
|
||||
return node
|
||||
}
|
||||
|
||||
func jsonString(payload []byte, path string) string {
|
||||
if s, ok := jsonWalk(payload, path).(string); ok {
|
||||
return s
|
||||
}
|
||||
return ""
|
||||
}
|
||||
|
||||
func jsonNumber(payload []byte, path string) int64 {
|
||||
if f, ok := jsonWalk(payload, path).(float64); ok {
|
||||
return int64(f)
|
||||
}
|
||||
return 0
|
||||
}
|
||||
Some files were not shown because too many files have changed in this diff Show More
Reference in New Issue
Block a user