diff --git a/common/singbridge/destination.go b/common/singbridge/destination.go deleted file mode 100644 index 217c4d082..000000000 --- a/common/singbridge/destination.go +++ /dev/null @@ -1,53 +0,0 @@ -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 -} diff --git a/common/singbridge/dialer.go b/common/singbridge/dialer.go deleted file mode 100644 index 07c428813..000000000 --- a/common/singbridge/dialer.go +++ /dev/null @@ -1,72 +0,0 @@ -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 -} diff --git a/common/singbridge/error.go b/common/singbridge/error.go deleted file mode 100644 index ac9e63517..000000000 --- a/common/singbridge/error.go +++ /dev/null @@ -1,10 +0,0 @@ -package singbridge - -import E "github.com/sagernet/sing/common/exceptions" - -func ReturnError(err error) error { - if E.IsClosedOrCanceled(err) { - return nil - } - return err -} diff --git a/common/singbridge/handler.go b/common/singbridge/handler.go deleted file mode 100644 index ee4b7c15f..000000000 --- a/common/singbridge/handler.go +++ /dev/null @@ -1,58 +0,0 @@ -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()) -} diff --git a/common/singbridge/logger.go b/common/singbridge/logger.go deleted file mode 100644 index 16ff29cc3..000000000 --- a/common/singbridge/logger.go +++ /dev/null @@ -1,70 +0,0 @@ -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) { -} diff --git a/common/singbridge/packet.go b/common/singbridge/packet.go deleted file mode 100644 index fde4bed3d..000000000 --- a/common/singbridge/packet.go +++ /dev/null @@ -1,107 +0,0 @@ -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 -} diff --git a/common/singbridge/pipe.go b/common/singbridge/pipe.go deleted file mode 100644 index 94c5ee0c1..000000000 --- a/common/singbridge/pipe.go +++ /dev/null @@ -1,81 +0,0 @@ -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 -} diff --git a/common/singbridge/reader.go b/common/singbridge/reader.go deleted file mode 100644 index 1ace1845f..000000000 --- a/common/singbridge/reader.go +++ /dev/null @@ -1,66 +0,0 @@ -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 -} diff --git a/go.mod b/go.mod index be24f79b3..097984025 100644 --- a/go.mod +++ b/go.mod @@ -18,8 +18,6 @@ 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 diff --git a/go.sum b/go.sum index f90697948..261a5736a 100644 --- a/go.sum +++ b/go.sum @@ -76,10 +76,6 @@ 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= diff --git a/infra/conf/shadowsocks.go b/infra/conf/shadowsocks.go index 18451ab5c..c4e7790e2 100644 --- a/infra/conf/shadowsocks.go +++ b/infra/conf/shadowsocks.go @@ -3,8 +3,6 @@ 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" @@ -55,7 +53,7 @@ func (v *ShadowsocksServerConfig) Build() (proto.Message, error) { v.Users = v.Clients } - if C.Contains(shadowaead_2022.List, v.Cipher) { + if shadowsocks_2022.IsSupportedMethod(v.Cipher) { return buildShadowsocks2022(v) } @@ -216,7 +214,7 @@ func (v *ShadowsocksClientConfig) Build() (proto.Message, error) { if len(v.Servers) == 1 { server := v.Servers[0] - if C.Contains(shadowaead_2022.List, server.Cipher) { + if shadowsocks_2022.IsSupportedMethod(server.Cipher) { if server.Address == nil { return nil, errors.New("Shadowsocks server address is not set.") } @@ -238,7 +236,7 @@ func (v *ShadowsocksClientConfig) Build() (proto.Message, error) { config := new(shadowsocks.ClientConfig) for _, server := range v.Servers { - if C.Contains(shadowaead_2022.List, server.Cipher) { + if shadowsocks_2022.IsSupportedMethod(server.Cipher) { return nil, errors.New("Shadowsocks 2022 accept no multi servers") } if server.Address == nil { diff --git a/proxy/shadowsocks_2022/cipher.go b/proxy/shadowsocks_2022/cipher.go new file mode 100644 index 000000000..dde381c92 --- /dev/null +++ b/proxy/shadowsocks_2022/cipher.go @@ -0,0 +1,59 @@ +package shadowsocks_2022 + +import ( + "crypto/aes" + "crypto/cipher" + "errors" + + "golang.org/x/crypto/chacha20poly1305" +) + +type CipherMethod struct { + Name string + KeySaltLength int + IsChaCha bool +} + +var ( + cipherAES128GCM = &CipherMethod{Name: MethodAES128GCM, KeySaltLength: 16, IsChaCha: false} + cipherAES256GCM = &CipherMethod{Name: MethodAES256GCM, KeySaltLength: 32, IsChaCha: false} + cipherChaCha20Poly1305 = &CipherMethod{Name: MethodChaCha20Poly1305, KeySaltLength: 32, IsChaCha: true} +) + +func GetCipherMethod(name string) (*CipherMethod, error) { + switch name { + case MethodAES128GCM: + return cipherAES128GCM, nil + case MethodAES256GCM: + return cipherAES256GCM, nil + case MethodChaCha20Poly1305: + return cipherChaCha20Poly1305, nil + default: + return nil, errors.New("unknown shadowsocks 2022 method") + } +} + +// NewAEAD creates standard stream AEAD cipher instance (AES-GCM or ChaCha20-Poly1305) +func (m *CipherMethod) NewAEAD(key []byte) (cipher.AEAD, error) { + if m.IsChaCha { + return chacha20poly1305.New(key) + } + block, err := aes.NewCipher(key) + if err != nil { + return nil, err + } + return cipher.NewGCM(block) +} + +// NewBlock creates standard 16-byte block cipher for AES header encryption/decryption +func (m *CipherMethod) NewBlock(key []byte) (cipher.Block, error) { + return aes.NewCipher(key) +} + +// NewUDPCipher creates AEAD cipher for UDP packets (XChaCha20-Poly1305 with 24-byte nonce) +func (m *CipherMethod) NewUDPCipher(key []byte) (cipher.AEAD, error) { + if m.IsChaCha { + return chacha20poly1305.NewX(key) + } + return nil, errors.New("shadowsocks-2022: udp separate AEAD cipher only available for chacha20 method") +} diff --git a/proxy/shadowsocks_2022/config.go b/proxy/shadowsocks_2022/config.go index 9ddd2cf88..2fa28aa0c 100644 --- a/proxy/shadowsocks_2022/config.go +++ b/proxy/shadowsocks_2022/config.go @@ -1,6 +1,9 @@ package shadowsocks_2022 import ( + "bytes" + "encoding/base64" + "google.golang.org/protobuf/proto" "github.com/xtls/xray-core/common/protocol" @@ -8,26 +11,31 @@ import ( // MemoryAccount is an account type converted from Account. type MemoryAccount struct { - Key string + Key []byte } // AsAccount implements protocol.AsAccount. func (u *Account) AsAccount() (protocol.Account, error) { + keyStr := u.GetKey() + raw, err := base64.StdEncoding.DecodeString(keyStr) + if err != nil { + raw = []byte(keyStr) + } return &MemoryAccount{ - Key: u.GetKey(), + Key: raw, }, 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 bytes.Equal(a.Key, account.Key) } return false } func (a *MemoryAccount) ToProto() proto.Message { return &Account{ - Key: a.Key, + Key: base64.StdEncoding.EncodeToString(a.Key), } } diff --git a/proxy/shadowsocks_2022/inbound.go b/proxy/shadowsocks_2022/inbound.go index edf9857c8..ca9ef9a5d 100644 --- a/proxy/shadowsocks_2022/inbound.go +++ b/proxy/shadowsocks_2022/inbound.go @@ -2,17 +2,11 @@ package shadowsocks_2022 import ( "context" + "io" "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/antireplay" "github.com/xtls/xray-core/common/buf" "github.com/xtls/xray-core/common/errors" "github.com/xtls/xray-core/common/log" @@ -20,7 +14,10 @@ import ( "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/task" + "github.com/xtls/xray-core/common/utils" + "github.com/xtls/xray-core/core" + "github.com/xtls/xray-core/features/policy" "github.com/xtls/xray-core/features/routing" "github.com/xtls/xray-core/transport/internet/stat" ) @@ -32,10 +29,13 @@ func init() { } type Inbound struct { - networks []net.Network - service shadowsocks.Service - email string - level int + networks []net.Network + method *CipherMethod + psk []byte + user *protocol.MemoryUser + saltFilter *antireplay.ReplayFilter[[32]byte] + udpCodec *UDPServerCodec + policyManager policy.Manager } func NewServer(ctx context.Context, config *ServerConfig) (*Inbound, error) { @@ -46,20 +46,35 @@ func NewServer(ctx context.Context, config *ServerConfig) (*Inbound, error) { 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) + + method, err := GetCipherMethod(config.Method) if err != nil { - return nil, errors.New("create service").Base(err) + return nil, errors.New("unsupported method: ", config.Method).Base(err) } - inbound.service = service - return inbound, nil + + psk, err := ParseKey(config.Key, method.KeySaltLength) + if err != nil { + return nil, err + } + + udpCodec, err := NewUDPServerCodec(method, psk, 500*time.Second) + if err != nil { + return nil, err + } + + v := core.MustFromContext(ctx) + return &Inbound{ + networks: networks, + method: method, + psk: psk, + saltFilter: antireplay.NewMapFilter[[32]byte](60), + user: &protocol.MemoryUser{ + Email: config.Email, + Level: uint32(config.Level), + }, + udpCodec: udpCodec, + policyManager: v.GetFeature(policy.ManagerType()).(policy.Manager), + }, nil } func (i *Inbound) Network() []net.Network { @@ -71,113 +86,193 @@ func (i *Inbound) Process(ctx context.Context, network net.Network, connection s inbound.Name = "shadowsocks-2022" inbound.CanSpliceCopy = 3 - var metadata M.Metadata - if inbound.Source.IsValid() { - metadata.Source = M.ParseSocksaddr(inbound.Source.NetAddr()) + if network == net.Network_TCP { + return i.processTCP(ctx, connection, dispatcher) + } + return i.processUDP(ctx, connection, dispatcher) +} + +func (i *Inbound) processTCP(ctx context.Context, conn net.Conn, dispatcher routing.Dispatcher) error { + defer conn.Close() + + sessionPolicy := i.policyManager.ForLevel(0) + if err := conn.SetReadDeadline(time.Now().Add(sessionPolicy.Timeouts.Handshake)); err != nil { + return errors.New("unable to set read deadline").Base(err).AtWarning() } - ctx = session.ContextWithDispatcher(ctx, dispatcher) + var salt [32]byte + saltSlice := salt[:i.method.KeySaltLength] + if _, err := io.ReadFull(conn, saltSlice); err != nil { + return err + } - 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 !i.saltFilter.Check(salt) { + return ErrSaltNotUnique + } + + sessionKey := DeriveSessionSubKey(i.psk, saltSlice, i.method.KeySaltLength) + aead, err := i.method.NewAEAD(sessionKey) + if err != nil { + return err + } + + reader := NewStreamReader(conn, aead) + + reqHeader, err := ReadClientRequestHeader(conn, reader) + if err != nil { + return err + } + _ = conn.SetReadDeadline(time.Time{}) + dest := reqHeader.Destination + + writer, err := WriteTCPResponse(conn, i.method, i.psk, saltSlice, nil) + if err != nil { + return err + } + + inbound := session.InboundFromContext(ctx) + if inbound == nil { + inbound = new(session.Inbound) + ctx = session.ContextWithInbound(ctx, inbound) + } + inbound.User = i.user + + ctx = log.ContextWithAccessMessage(ctx, &log.AccessMessage{ + From: conn.RemoteAddr(), + To: dest, + Status: log.AccessAccepted, + Email: i.user.Email, + }) + + errors.LogInfo(ctx, "tunneling request to ", dest) + + link, err := dispatcher.Dispatch(ctx, dest) + if err != nil { + return err + } + + if len(reqHeader.EarlyData) > 0 { + earlyBuf := buf.New() + earlyBuf.Write(reqHeader.EarlyData) + if err := link.Writer.WriteMultiBuffer(buf.MultiBuffer{earlyBuf}); err != nil { + return err + } + } + + sessionPolicy = i.policyManager.ForLevel(uint32(i.user.Level)) + ctx, cancel := context.WithCancel(ctx) + timer := signal.CancelAfterInactivity(ctx, cancel, sessionPolicy.Timeouts.ConnectionIdle) + ctx = policy.ContextWithBufferPolicy(ctx, sessionPolicy.Buffer) + + requestDone := func() error { + defer timer.SetTimeout(sessionPolicy.Timeouts.DownlinkOnly) + return buf.Copy(reader, link.Writer, buf.UpdateActivity(timer)) + } + + responseDone := func() error { + defer timer.SetTimeout(sessionPolicy.Timeouts.UplinkOnly) + return buf.Copy(link.Reader, writer, buf.UpdateActivity(timer)) + } + + responseDoneAndCloseWriter := task.OnSuccess(responseDone, task.Close(link.Writer)) + return task.Run(ctx, requestDone, responseDoneAndCloseWriter) +} + +func (i *Inbound) processUDP(ctx context.Context, conn stat.Connection, dispatcher routing.Dispatcher) error { + udpConns := utils.NewTypedSyncMap[uint64, *udpConnEntry]() + defer func() { + udpConns.Range(func(key uint64, entry *udpConnEntry) bool { + entry.timer.SetTimeout(0) + return true + }) + }() + + inbound := session.InboundFromContext(ctx) + if inbound == nil { + inbound = new(session.Inbound) + ctx = session.ContextWithInbound(ctx, inbound) + } + inbound.User = i.user + + reader := buf.NewReader(conn) + for { + mb, err := reader.ReadMultiBuffer() + if err != nil { + buf.ReleaseMulti(mb) + return err + } + + for _, b := range mb { + decoded, err := i.udpCodec.DecodePacket(b.Bytes()) if err != nil { - buf.ReleaseMulti(mb) - return singbridge.ReturnError(err) + b.Release() + continue } - for _, buffer := range mb { - packet := B.As(buffer.Bytes()).ToOwned() - buffer.Release() - err = i.service.NewPacket(ctx, pc, packet, metadata) + + entry, ok := udpConns.Load(decoded.SessionID) + if !ok { + sessCtx, cancel := context.WithCancel(ctx) + sessCtx = log.ContextWithAccessMessage(sessCtx, &log.AccessMessage{ + From: conn.RemoteAddr(), + To: decoded.Destination, + Status: log.AccessAccepted, + Email: i.user.Email, + }) + + link, err := dispatcher.Dispatch(sessCtx, decoded.Destination) if err != nil { - packet.Release() - buf.ReleaseMulti(mb) - return err + cancel() + b.Release() + continue + } + + newEntry := &udpConnEntry{ + link: link, + cancel: cancel, + } + sessionPolicy := i.policyManager.ForLevel(uint32(i.user.Level)) + newEntry.timer = signal.CancelAfterInactivity(sessCtx, func() { + udpConns.Delete(decoded.SessionID) + common.Interrupt(link.Reader) + common.Interrupt(link.Writer) + cancel() + }, sessionPolicy.Timeouts.ConnectionIdle) + + actual, loaded := udpConns.LoadOrStore(decoded.SessionID, newEntry) + if loaded { + // Another goroutine/packet beat us to storing, terminate our redundant link + newEntry.timer.SetTimeout(0) + entry = actual + } else { + entry = newEntry + go func(sessID uint64, dest net.Destination, cEntry *udpConnEntry) { + defer func() { + cEntry.timer.SetTimeout(0) + }() + for { + resMb, err := cEntry.link.Reader.ReadMultiBuffer() + if err != nil { + return + } + cEntry.timer.Update() + for _, rb := range resMb { + encPacket, err := i.udpCodec.EncodePacket(sessID, dest, rb.Bytes()) + rb.Release() + if err != nil { + continue + } + _, _ = conn.Write(encPacket) + } + } + }(decoded.SessionID, decoded.Destination, entry) } } + + entry.timer.Update() + payloadBuf := buf.New() + payloadBuf.Write(decoded.Payload) + b.Release() + _ = entry.link.Writer.WriteMultiBuffer(buf.MultiBuffer{payloadBuf}) } } } - -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 -} diff --git a/proxy/shadowsocks_2022/inbound_multi.go b/proxy/shadowsocks_2022/inbound_multi.go index d6d68c09a..7de2ab8a7 100644 --- a/proxy/shadowsocks_2022/inbound_multi.go +++ b/proxy/shadowsocks_2022/inbound_multi.go @@ -2,21 +2,17 @@ package shadowsocks_2022 import ( "context" - "encoding/base64" + "crypto/cipher" + "encoding/binary" + "io" "strconv" "strings" "sync" + "sync/atomic" "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/antireplay" "github.com/xtls/xray-core/common/buf" "github.com/xtls/xray-core/common/errors" "github.com/xtls/xray-core/common/log" @@ -24,8 +20,11 @@ import ( "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/task" + "github.com/xtls/xray-core/common/utils" "github.com/xtls/xray-core/common/uuid" + "github.com/xtls/xray-core/core" + "github.com/xtls/xray-core/features/policy" "github.com/xtls/xray-core/features/routing" "github.com/xtls/xray-core/transport/internet/stat" ) @@ -38,9 +37,16 @@ func init() { type MultiUserInbound struct { sync.Mutex - networks []net.Network - users []*protocol.MemoryUser - service *shadowaead_2022.MultiService[int] + networks []net.Network + method *CipherMethod + masterPSK []byte + usersByHash *utils.TypedSyncMap[[AESBlockSize]byte, *protocol.MemoryUser] + usersByEmail *utils.TypedSyncMap[string, *protocol.MemoryUser] + userCount atomic.Int64 + saltFilter *antireplay.ReplayFilter[[32]byte] + udpSessions *UDPSessionManager + udpMasterCipher cipher.Block + policyManager policy.Manager } func NewMultiServer(ctx context.Context, config *MultiUserServerConfig) (*MultiUserInbound, error) { @@ -51,138 +57,131 @@ func NewMultiServer(ctx context.Context, config *MultiUserServerConfig) (*MultiU net.Network_UDP, } } - memUsers := []*protocol.MemoryUser{} - for i, user := range config.Users { + + method, err := GetCipherMethod(config.Method) + if err != nil { + return nil, err + } + if method.IsChaCha { + return nil, errors.New("shadowsocks 2022 multi-user: only aes methods are supported") + } + + masterPSK, err := ParseKey(config.Key, method.KeySaltLength) + if err != nil { + return nil, err + } + + masterBlock, err := method.NewBlock(masterPSK) + if err != nil { + return nil, err + } + + v := core.MustFromContext(ctx) + i := &MultiUserInbound{ + networks: networks, + method: method, + masterPSK: masterPSK, + usersByHash: utils.NewTypedSyncMap[[AESBlockSize]byte, *protocol.MemoryUser](), + usersByEmail: utils.NewTypedSyncMap[string, *protocol.MemoryUser](), + saltFilter: antireplay.NewMapFilter[[32]byte](60), + udpSessions: NewUDPSessionManager(500 * time.Second), + udpMasterCipher: masterBlock, + policyManager: v.GetFeature(policy.ManagerType()).(policy.Manager), + } + + for idx, user := range config.Users { if user.Email == "" { u := uuid.New() - user.Email = "unnamed-user-" + strconv.Itoa(i) + "-" + u.String() + user.Email = "unnamed-user-" + strconv.Itoa(idx) + "-" + u.String() } - u, err := user.ToMemoryUser() + memUser, err := user.ToMemoryUser() if err != nil { - return nil, errors.New("failed to get shadowsocks user").Base(err).AtError() + return nil, errors.New("failed to parse shadowsocks user").Base(err) + } + if err := i.AddUser(ctx, memUser); err != nil { + return nil, err } - 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 + return i, nil } -// AddUser implements proxy.UserManager.AddUser(). +// AddUser implements proxy.UserManager.AddUser() func (i *MultiUserInbound) AddUser(ctx context.Context, u *protocol.MemoryUser) error { i.Lock() defer i.Unlock() + var emailKey string if u.Email != "" { - for idx := range i.users { - if i.users[idx].Email == u.Email { - return errors.New("User ", u.Email, " already exists.") - } + emailKey = strings.ToLower(u.Email) + if _, exists := i.usersByEmail.Load(emailKey); exists { + 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 }), - ) + memAcc, ok := u.Account.(*MemoryAccount) + if !ok { + return errors.New("missing or invalid user account") + } + + if len(memAcc.Key) != i.method.KeySaltLength { + return ErrBadKey + } + + pskHash := DeriveUserPSKHash(memAcc.Key) + i.usersByHash.Store(pskHash, u) + if emailKey != "" { + i.usersByEmail.Store(emailKey, u) + } + i.userCount.Add(1) return nil } -// RemoveUser implements proxy.UserManager.RemoveUser(). +// 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.") + 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 - } + emailKey := strings.ToLower(email) + u, loaded := i.usersByEmail.LoadAndDelete(emailKey) + if !loaded { + return errors.New("user ", email, " not found") } - 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 }), - ) + pskHash := DeriveUserPSKHash(u.Account.(*MemoryAccount).Key) + i.usersByHash.Delete(pskHash) + i.userCount.Add(-1) return nil } -// GetUser implements proxy.UserManager.GetUser(). +// 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 + u, _ := i.usersByEmail.Load(strings.ToLower(email)) + return u } -// GetUsers implements proxy.UserManager.GetUsers(). +// 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 + var users []*protocol.MemoryUser + i.usersByEmail.Range(func(_ string, user *protocol.MemoryUser) bool { + users = append(users, user) + return true + }) + return users } -// GetUsersCount implements proxy.UserManager.GetUsersCount(). +// GetUsersCount implements proxy.UserManager.GetUsersCount() func (i *MultiUserInbound) GetUsersCount(context.Context) int64 { - i.Lock() - defer i.Unlock() - return int64(len(i.users)) + return i.userCount.Load() } func (i *MultiUserInbound) Network() []net.Network { @@ -194,97 +193,327 @@ func (i *MultiUserInbound) Process(ctx context.Context, network net.Network, con inbound.Name = "shadowsocks-2022-multi" inbound.CanSpliceCopy = 3 - var metadata M.Metadata - if inbound.Source.IsValid() { - metadata.Source = M.ParseSocksaddr(inbound.Source.NetAddr()) + if network == net.Network_TCP { + return i.processTCP(ctx, connection, dispatcher) + } + return i.processUDP(ctx, connection, dispatcher) +} + +func (i *MultiUserInbound) processTCP(ctx context.Context, conn net.Conn, dispatcher routing.Dispatcher) error { + defer conn.Close() + + sessionPolicy := i.policyManager.ForLevel(0) + if err := conn.SetReadDeadline(time.Now().Add(sessionPolicy.Timeouts.Handshake)); err != nil { + return errors.New("unable to set read deadline").Base(err).AtWarning() } - ctx = session.ContextWithDispatcher(ctx, dispatcher) + // 1. Read Request Salt (16 or 32 bytes) + var salt [32]byte + saltSlice := salt[:i.method.KeySaltLength] + if _, err := io.ReadFull(conn, saltSlice); err != nil { + return err + } - 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) + if !i.saltFilter.Check(salt) { + return ErrSaltNotUnique + } + + // 2. Read Extended Identity Header (16 bytes) + var eih [AESBlockSize]byte + if _, err := io.ReadFull(conn, eih[:]); err != nil { + return err + } + + // Decrypt EIH with IdentitySubKey derived from masterPSK and salt + identitySubkey := DeriveIdentitySubKey(i.masterPSK, saltSlice, i.method.KeySaltLength) + block, err := i.method.NewBlock(identitySubkey) + if err != nil { + return err + } + + var decryptedHash [AESBlockSize]byte + block.Decrypt(decryptedHash[:], eih[:]) + + // Lookup user + user, ok := i.usersByHash.Load(decryptedHash) + if !ok || user == nil { + return ErrInvalidRequest + } + userPSK := user.Account.(*MemoryAccount).Key + + // 3. Derive Session Subkey using matched user's PSK + sessionKey := DeriveSessionSubKey(userPSK, saltSlice, i.method.KeySaltLength) + aead, err := i.method.NewAEAD(sessionKey) + if err != nil { + return err + } + + reader := NewStreamReader(conn, aead) + + // 4 & 5. Read Client Request Header + reqHeader, err := ReadClientRequestHeader(conn, reader) + if err != nil { + return err + } + _ = conn.SetReadDeadline(time.Time{}) + dest := reqHeader.Destination + + // 6. Send Server Response Handshake + writer, err := WriteTCPResponse(conn, i.method, userPSK, saltSlice, nil) + if err != nil { + return err + } + + // 7. Dispatch Connection to Xray routing with matched User + inbound := session.InboundFromContext(ctx) + if inbound == nil { + inbound = new(session.Inbound) + ctx = session.ContextWithInbound(ctx, inbound) + } + inbound.User = user + + ctx = log.ContextWithAccessMessage(ctx, &log.AccessMessage{ + From: conn.RemoteAddr(), + To: dest, + Status: log.AccessAccepted, + Email: user.Email, + }) + + errors.LogInfo(ctx, "tunneling request to ", dest, " for user ", user.Email) + + link, err := dispatcher.Dispatch(ctx, dest) + if err != nil { + return err + } + + if len(reqHeader.EarlyData) > 0 { + earlyBuf := buf.New() + earlyBuf.Write(reqHeader.EarlyData) + if err := link.Writer.WriteMultiBuffer(buf.MultiBuffer{earlyBuf}); err != nil { + return err + } + } + + sessionPolicy = i.policyManager.ForLevel(user.Level) + ctx, cancel := context.WithCancel(ctx) + timer := signal.CancelAfterInactivity(ctx, cancel, sessionPolicy.Timeouts.ConnectionIdle) + ctx = policy.ContextWithBufferPolicy(ctx, sessionPolicy.Buffer) + + requestDone := func() error { + defer timer.SetTimeout(sessionPolicy.Timeouts.DownlinkOnly) + return buf.Copy(reader, link.Writer, buf.UpdateActivity(timer)) + } + + responseDone := func() error { + defer timer.SetTimeout(sessionPolicy.Timeouts.UplinkOnly) + return buf.Copy(link.Reader, writer, buf.UpdateActivity(timer)) + } + + responseDoneAndCloseWriter := task.OnSuccess(responseDone, task.Close(link.Writer)) + return task.Run(ctx, requestDone, responseDoneAndCloseWriter) +} + +func (i *MultiUserInbound) processUDP(ctx context.Context, conn stat.Connection, dispatcher routing.Dispatcher) error { + udpConns := utils.NewTypedSyncMap[uint64, *udpConnEntry]() + defer func() { + udpConns.Range(func(key uint64, entry *udpConnEntry) bool { + entry.timer.SetTimeout(0) + return true + }) + }() + + reader := buf.NewReader(conn) + for { + mb, err := reader.ReadMultiBuffer() + if err != nil { + buf.ReleaseMulti(mb) + return err + } + + for _, b := range mb { + // In multi-user UDP: + // Packet header is 16 bytes: Encrypted(SessionID + PacketID) + // Followed by 16 bytes EIH + packetBytes := b.Bytes() + if len(packetBytes) < 32+1+8+2 { + b.Release() + continue } - for _, buffer := range mb { - packet := B.As(buffer.Bytes()).ToOwned() - buffer.Release() - err = i.service.NewPacket(ctx, pc, packet, metadata) + + var rawHeader [16]byte + i.udpMasterCipher.Decrypt(rawHeader[:], packetBytes[:16]) + + sessionID := binary.BigEndian.Uint64(rawHeader[:8]) + packetID := binary.BigEndian.Uint64(rawHeader[8:16]) + + // Replay protection & session lookup + sessionItem, _ := i.udpSessions.GetOrCreate(sessionID) + + sessionItem.Lock() + if !sessionItem.Window.Check(packetID) { + sessionItem.Unlock() + b.Release() + continue + } + + var userPSK []byte + var currentUser *protocol.MemoryUser + if sessionItem.User != nil { + currentUser = sessionItem.User + userPSK = sessionItem.UserPSK + sessionItem.Unlock() + } else { + sessionItem.Unlock() + // Decrypt EIH + identitySubkey := DeriveIdentitySubKey(i.masterPSK, rawHeader[:8], i.method.KeySaltLength) + idBlock, err := i.method.NewBlock(identitySubkey) if err != nil { - packet.Release() - buf.ReleaseMulti(mb) - return err + b.Release() + continue + } + + var decryptedHash [16]byte + idBlock.Decrypt(decryptedHash[:], packetBytes[16:32]) + + user, ok := i.usersByHash.Load(decryptedHash) + if !ok || user == nil { + b.Release() + continue + } + currentUser = user + userPSK = user.Account.(*MemoryAccount).Key + + sessionItem.Lock() + sessionItem.User = user + sessionItem.UserPSK = userPSK + sessionItem.Unlock() + } + + // Decrypt Body (with AEAD caching per session) + bodyAead := sessionItem.GetRemoteCipher() + if bodyAead == nil { + bodyKey := DeriveSessionSubKey(userPSK, rawHeader[:8], i.method.KeySaltLength) + var err error + bodyAead, err = i.method.NewAEAD(bodyKey) + if err != nil { + b.Release() + continue + } + sessionItem.SetRemoteCipher(bodyAead) + } + + bodyNonce := rawHeader[4:16] + bodyCipher := packetBytes[32:] + bodyPlain, err := bodyAead.Open(nil, bodyNonce, bodyCipher, nil) + b.Release() + if err != nil || len(bodyPlain) < 1+8+2 { + continue + } + + sessionItem.Lock() + sessionItem.Window.Add(packetID) + sessionItem.Unlock() + + if bodyPlain[0] != HeaderTypeClient { + continue + } + epoch := binary.BigEndian.Uint64(bodyPlain[1:9]) + diff := time.Now().Unix() - int64(epoch) + if diff < -30 || diff > 30 { + continue + } + + paddingLen := int(binary.BigEndian.Uint16(bodyPlain[9:11])) + offset := 11 + paddingLen + if len(bodyPlain) < offset { + continue + } + + dest, addrLen, err := parseAddressPort(bodyPlain[offset:]) + if err != nil { + continue + } + + payload := bodyPlain[offset+addrLen:] + payloadCopy := make([]byte, len(payload)) + copy(payloadCopy, payload) + + entry, ok := udpConns.Load(sessionID) + if !ok { + sessCtx, cancel := context.WithCancel(ctx) + inbound := session.InboundFromContext(sessCtx) + if inbound == nil { + inbound = new(session.Inbound) + sessCtx = session.ContextWithInbound(sessCtx, inbound) + } + inbound.User = currentUser + + sessCtx = log.ContextWithAccessMessage(sessCtx, &log.AccessMessage{ + From: conn.RemoteAddr(), + To: dest, + Status: log.AccessAccepted, + Email: currentUser.Email, + }) + + link, err := dispatcher.Dispatch(sessCtx, dest) + if err != nil { + cancel() + continue + } + + newEntry := &udpConnEntry{ + link: link, + cancel: cancel, + } + sessionPolicy := i.policyManager.ForLevel(currentUser.Level) + newEntry.timer = signal.CancelAfterInactivity(sessCtx, func() { + udpConns.Delete(sessionID) + common.Interrupt(link.Reader) + common.Interrupt(link.Writer) + cancel() + }, sessionPolicy.Timeouts.ConnectionIdle) + + actual, loaded := udpConns.LoadOrStore(sessionID, newEntry) + if loaded { + newEntry.timer.SetTimeout(0) + entry = actual + } else { + entry = newEntry + go func(sessID uint64, uPSK []byte, d net.Destination, cEntry *udpConnEntry) { + defer func() { + cEntry.timer.SetTimeout(0) + }() + for { + resMb, err := cEntry.link.Reader.ReadMultiBuffer() + if err != nil { + return + } + cEntry.timer.Update() + for _, rb := range resMb { + encPacket, err := i.encodeServerUDPPacket(sessID, uPSK, d, rb.Bytes()) + rb.Release() + if err != nil { + continue + } + _, _ = conn.Write(encPacket) + } + } + }(sessionID, userPSK, dest, entry) } } + + entry.timer.Update() + pBuf := buf.New() + pBuf.Write(payloadCopy) + _ = entry.link.Writer.WriteMultiBuffer(buf.MultiBuffer{pBuf}) } } } -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 +func (i *MultiUserInbound) encodeServerUDPPacket(clientSessionID uint64, userPSK []byte, dest net.Destination, payload []byte) ([]byte, error) { + sessionItem, _ := i.udpSessions.GetOrCreate(clientSessionID) + if err := sessionItem.EnsureServerState(i.method, i.udpMasterCipher, nil, userPSK); err != nil { + return nil, 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()) + return sessionItem.EncodeServerPacket(i.method, clientSessionID, dest, payload) } diff --git a/proxy/shadowsocks_2022/inbound_relay.go b/proxy/shadowsocks_2022/inbound_relay.go index 4ca5e2075..625ba19a2 100644 --- a/proxy/shadowsocks_2022/inbound_relay.go +++ b/proxy/shadowsocks_2022/inbound_relay.go @@ -2,18 +2,12 @@ package shadowsocks_2022 import ( "context" + "crypto/cipher" + "encoding/binary" + "io" "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" @@ -22,8 +16,11 @@ import ( "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/task" + "github.com/xtls/xray-core/common/utils" "github.com/xtls/xray-core/common/uuid" + "github.com/xtls/xray-core/core" + "github.com/xtls/xray-core/features/policy" "github.com/xtls/xray-core/features/routing" "github.com/xtls/xray-core/transport/internet/stat" ) @@ -34,10 +31,22 @@ func init() { })) } +type relayDest struct { + destination net.Destination + email string + level uint32 + key []byte + blockCipher cipher.Block +} + type RelayInbound struct { - networks []net.Network - destinations []*RelayDestination - service *shadowaead_2022.RelayService[int] + networks []net.Network + method *CipherMethod + relayPSK []byte + relayBlock cipher.Block + destinations map[[AESBlockSize]byte]*relayDest + rawDestinations []*RelayDestination + policyManager policy.Manager } func NewRelayServer(ctx context.Context, config *RelayServerConfig) (*RelayInbound, error) { @@ -48,39 +57,63 @@ func NewRelayServer(ctx context.Context, config *RelayServerConfig) (*RelayInbou 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) + + method, err := GetCipherMethod(config.Method) if err != nil { - return nil, errors.New("create service").Base(err) + return nil, err + } + if method.IsChaCha { + return nil, errors.New("shadowsocks 2022 relay: only aes methods are supported") } - for i, destination := range config.Destinations { - if destination.Email == "" { + relayPSK, err := ParseKey(config.Key, method.KeySaltLength) + if err != nil { + return nil, err + } + + relayBlock, err := method.NewBlock(relayPSK) + if err != nil { + return nil, err + } + + v := core.MustFromContext(ctx) + i := &RelayInbound{ + networks: networks, + method: method, + relayPSK: relayPSK, + relayBlock: relayBlock, + destinations: make(map[[AESBlockSize]byte]*relayDest), + rawDestinations: config.Destinations, + policyManager: v.GetFeature(policy.ManagerType()).(policy.Manager), + } + + for idx, d := range config.Destinations { + if d.Email == "" { u := uuid.New() - destination.Email = "unnamed-destination-" + strconv.Itoa(i) + "-" + u.String() + d.Email = "unnamed-destination-" + strconv.Itoa(idx) + "-" + u.String() + } + destKey, err := ParseKey(d.Key, method.KeySaltLength) + if err != nil { + return nil, err + } + + destBlock, err := method.NewBlock(destKey) + if err != nil { + return nil, err + } + + hash := DeriveUserPSKHash(destKey) + + i.destinations[hash] = &relayDest{ + destination: net.TCPDestination(d.Address.AsAddress(), net.Port(d.Port)), + email: d.Email, + level: uint32(d.Level), + key: destKey, + blockCipher: destBlock, } } - 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 + + return i, nil } func (i *RelayInbound) Network() []net.Network { @@ -92,103 +125,214 @@ func (i *RelayInbound) Process(ctx context.Context, network net.Network, connect inbound.Name = "shadowsocks-2022-relay" inbound.CanSpliceCopy = 3 - var metadata M.Metadata - if inbound.Source.IsValid() { - metadata.Source = M.ParseSocksaddr(inbound.Source.NetAddr()) + if network == net.Network_TCP { + return i.processTCP(ctx, connection, dispatcher) + } + return i.processUDP(ctx, connection, dispatcher) +} + +func (i *RelayInbound) processTCP(ctx context.Context, conn net.Conn, dispatcher routing.Dispatcher) error { + defer conn.Close() + + sessionPolicy := i.policyManager.ForLevel(0) + if err := conn.SetReadDeadline(time.Now().Add(sessionPolicy.Timeouts.Handshake)); err != nil { + return errors.New("unable to set read deadline").Base(err).AtWarning() } - ctx = session.ContextWithDispatcher(ctx, dispatcher) + // Read Salt + Outer EIH + needed := i.method.KeySaltLength + AESBlockSize + var headerBuf [48]byte + headerSlice := headerBuf[:needed] + if _, err := io.ReadFull(conn, headerSlice); err != nil { + return err + } - 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) + salt := headerSlice[:i.method.KeySaltLength] + eih := headerSlice[i.method.KeySaltLength:] + + identitySubkey := DeriveIdentitySubKey(i.relayPSK, salt, i.method.KeySaltLength) + block, err := i.method.NewBlock(identitySubkey) + if err != nil { + return err + } + + var decryptedHash [AESBlockSize]byte + block.Decrypt(decryptedHash[:], eih) + + targetDest, ok := i.destinations[decryptedHash] + if !ok { + return ErrInvalidRequest + } + _ = conn.SetReadDeadline(time.Time{}) + + inbound := session.InboundFromContext(ctx) + if inbound == nil { + inbound = new(session.Inbound) + ctx = session.ContextWithInbound(ctx, inbound) + } + inbound.User = &protocol.MemoryUser{ + Email: targetDest.email, + Level: targetDest.level, + } + + ctx = log.ContextWithAccessMessage(ctx, &log.AccessMessage{ + From: conn.RemoteAddr(), + To: targetDest.destination, + Status: log.AccessAccepted, + Email: targetDest.email, + }) + + errors.LogInfo(ctx, "relaying connection to ", targetDest.destination) + + link, err := dispatcher.Dispatch(ctx, targetDest.destination) + if err != nil { + return err + } + + // Unwrap outer EIH: send client salt to next hop, stripping this hop's EIH + saltBuf := buf.New() + saltBuf.Write(salt) + if err := link.Writer.WriteMultiBuffer(buf.MultiBuffer{saltBuf}); err != nil { + return err + } + + sessionPolicy = i.policyManager.ForLevel(targetDest.level) + ctx, cancel := context.WithCancel(ctx) + timer := signal.CancelAfterInactivity(ctx, cancel, sessionPolicy.Timeouts.ConnectionIdle) + ctx = policy.ContextWithBufferPolicy(ctx, sessionPolicy.Buffer) + + requestDone := func() error { + defer timer.SetTimeout(sessionPolicy.Timeouts.DownlinkOnly) + return buf.Copy(buf.NewReader(conn), link.Writer, buf.UpdateActivity(timer)) + } + + responseDone := func() error { + defer timer.SetTimeout(sessionPolicy.Timeouts.UplinkOnly) + return buf.Copy(link.Reader, buf.NewWriter(conn), buf.UpdateActivity(timer)) + } + + responseDoneAndCloseWriter := task.OnSuccess(responseDone, task.Close(link.Writer)) + return task.Run(ctx, requestDone, responseDoneAndCloseWriter) +} + +func (i *RelayInbound) processUDP(ctx context.Context, conn stat.Connection, dispatcher routing.Dispatcher) error { + udpConns := utils.NewTypedSyncMap[uint64, *udpConnEntry]() + defer func() { + udpConns.Range(func(key uint64, entry *udpConnEntry) bool { + entry.timer.SetTimeout(0) + return true + }) + }() + + reader := buf.NewReader(conn) + for { + mb, err := reader.ReadMultiBuffer() + if err != nil { + buf.ReleaseMulti(mb) + return err + } + + for _, b := range mb { + data := b.Bytes() + if len(data) < 2*AESBlockSize { + b.Release() + continue } - for _, buffer := range mb { - packet := B.As(buffer.Bytes()).ToOwned() - buffer.Release() - err = i.service.NewPacket(ctx, pc, packet, metadata) + + var packetHeader [AESBlockSize]byte + i.relayBlock.Decrypt(packetHeader[:], data[:AESBlockSize]) + + var eiHeader [AESBlockSize]byte + i.relayBlock.Decrypt(eiHeader[:], data[AESBlockSize:2*AESBlockSize]) + for idx := 0; idx < AESBlockSize; idx++ { + eiHeader[idx] ^= packetHeader[idx] + } + + targetDest, ok := i.destinations[eiHeader] + if !ok { + b.Release() + continue + } + + // Extract sessionID from raw packetHeader for session-level link caching before re-encrypting + sessionID := binary.BigEndian.Uint64(packetHeader[:8]) + + // Re-encrypt packetHeader with next hop block cipher + targetDest.blockCipher.Encrypt(packetHeader[:], packetHeader[:]) + + // Strip outer EIH: replace second block with re-encrypted packetHeader and advance + copy(data[AESBlockSize:2*AESBlockSize], packetHeader[:]) + b.Advance(int32(AESBlockSize)) + + dest := targetDest.destination + dest.Network = net.Network_UDP + + entry, ok := udpConns.Load(sessionID) + if !ok { + sessCtx, cancel := context.WithCancel(ctx) + inbound := session.InboundFromContext(sessCtx) + if inbound == nil { + inbound = new(session.Inbound) + sessCtx = session.ContextWithInbound(sessCtx, inbound) + } + inbound.User = &protocol.MemoryUser{ + Email: targetDest.email, + Level: targetDest.level, + } + + sessCtx = log.ContextWithAccessMessage(sessCtx, &log.AccessMessage{ + From: conn.RemoteAddr(), + To: dest, + Status: log.AccessAccepted, + Email: targetDest.email, + }) + + link, err := dispatcher.Dispatch(sessCtx, dest) if err != nil { - packet.Release() - buf.ReleaseMulti(mb) - return err + cancel() + b.Release() + continue + } + + newEntry := &udpConnEntry{ + link: link, + cancel: cancel, + } + sessionPolicy := i.policyManager.ForLevel(targetDest.level) + newEntry.timer = signal.CancelAfterInactivity(sessCtx, func() { + udpConns.Delete(sessionID) + common.Interrupt(link.Reader) + common.Interrupt(link.Writer) + cancel() + }, sessionPolicy.Timeouts.ConnectionIdle) + + actual, loaded := udpConns.LoadOrStore(sessionID, newEntry) + if loaded { + newEntry.timer.SetTimeout(0) + entry = actual + } else { + entry = newEntry + go func(cEntry *udpConnEntry) { + defer func() { + cEntry.timer.SetTimeout(0) + }() + for { + resMb, err := cEntry.link.Reader.ReadMultiBuffer() + if err != nil { + return + } + cEntry.timer.Update() + for _, rb := range resMb { + _, _ = conn.Write(rb.Bytes()) + rb.Release() + } + } + }(entry) } } + + entry.timer.Update() + _ = entry.link.Writer.WriteMultiBuffer(buf.MultiBuffer{b}) } } } - -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()) -} diff --git a/proxy/shadowsocks_2022/kdf.go b/proxy/shadowsocks_2022/kdf.go new file mode 100644 index 000000000..3ebc18379 --- /dev/null +++ b/proxy/shadowsocks_2022/kdf.go @@ -0,0 +1,63 @@ +package shadowsocks_2022 + +import ( + "encoding/base64" + "strings" + + "lukechampine.com/blake3" +) + +const ( + ContextSessionSubKey = "shadowsocks 2022 session subkey" + ContextIdentitySubKey = "shadowsocks 2022 identity subkey" +) + +// ParseKey decodes a base64 or raw PSK key string and validates its length +func ParseKey(key string, keyLength int) ([]byte, error) { + raw, err := base64.StdEncoding.DecodeString(key) + if err != nil { + raw = []byte(key) + } + if len(raw) != keyLength { + return nil, ErrBadKey + } + return raw, nil +} + +func ParsePSKList(password string, keyLength int) ([][]byte, error) { + parts := strings.Split(password, ":") + pskList := make([][]byte, len(parts)) + for i, part := range parts { + norm, err := ParseKey(part, keyLength) + if err != nil { + return nil, err + } + pskList[i] = norm + } + return pskList, nil +} + +func deriveSubKey(ctx string, psk, salt []byte, keyLength int) []byte { + var keyMaterial [64]byte + kmLen := len(psk) + len(salt) + copy(keyMaterial[:], psk) + copy(keyMaterial[len(psk):], salt) + out := make([]byte, keyLength) + blake3.DeriveKey(out, ctx, keyMaterial[:kmLen]) + return out +} + +func DeriveSessionSubKey(psk, salt []byte, keyLength int) []byte { + return deriveSubKey(ContextSessionSubKey, psk, salt, keyLength) +} + +func DeriveIdentitySubKey(psk, salt []byte, keyLength int) []byte { + return deriveSubKey(ContextIdentitySubKey, psk, salt, keyLength) +} + +func DeriveUserPSKHash(userPSK []byte) [AESBlockSize]byte { + h := blake3.Sum512(userPSK) + var out [AESBlockSize]byte + copy(out[:], h[:AESBlockSize]) + return out +} diff --git a/proxy/shadowsocks_2022/outbound.go b/proxy/shadowsocks_2022/outbound.go index 5d1b9c9fb..f42590d81 100644 --- a/proxy/shadowsocks_2022/outbound.go +++ b/proxy/shadowsocks_2022/outbound.go @@ -2,21 +2,20 @@ package shadowsocks_2022 import ( "context" + "crypto/rand" + "io" "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/retry" "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/task" + "github.com/xtls/xray-core/core" + "github.com/xtls/xray-core/features/policy" "github.com/xtls/xray-core/transport" "github.com/xtls/xray-core/transport/internet" ) @@ -28,42 +27,49 @@ func init() { } type Outbound struct { - ctx context.Context - server net.Destination - method shadowsocks.Method + ctx context.Context + server net.Destination + method *CipherMethod + pskList [][]byte + finalPSK []byte + udpCodec *UDPPacketCodec + policyManager policy.Manager } func NewClient(ctx context.Context, config *ClientConfig) (*Outbound, error) { - o := &Outbound{ + method, err := GetCipherMethod(config.Method) + if err != nil { + return nil, errors.New("unsupported method: ", config.Method).Base(err) + } + + pskList, err := ParsePSKList(config.Key, method.KeySaltLength) + if err != nil { + return nil, errors.New("invalid key: ", config.Key).Base(err) + } + + finalPSK := pskList[len(pskList)-1] + udpCodec, err := NewUDPPacketCodec(method, finalPSK) + if err != nil { + return nil, errors.New("failed to create udp packet codec").Base(err) + } + + v := core.MustFromContext(ctx) + return &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 + method: method, + pskList: pskList, + finalPSK: finalPSK, + udpCodec: udpCodec, + policyManager: v.GetFeature(policy.ManagerType()).(policy.Manager), + }, 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() { @@ -78,70 +84,123 @@ func (o *Outbound) Process(ctx context.Context, link *transport.Link, dialer int 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() + var conn net.Conn + if err := retry.ExponentialBackoff(5, 100).On(func() error { + rawConn, err := dialer.Dial(ctx, serverDestination) + if err != nil { + return err + } + conn = rawConn + return nil + }); err != nil { + return errors.New("failed to find an available destination").Base(err).AtWarning() + } + defer conn.Close() + + var newCtx context.Context + var newCancel context.CancelFunc if session.TimeoutOnlyFromContext(ctx) { - ctx, _ = context.WithCancel(context.Background()) + newCtx, newCancel = context.WithCancel(context.Background()) + } + + sessionPolicy := o.policyManager.ForLevel(0) + ctx, cancel := context.WithCancel(ctx) + timer := signal.CancelAfterInactivity(ctx, func() { + cancel() + if newCancel != nil { + newCancel() + } + }, sessionPolicy.Timeouts.ConnectionIdle) + + ctx = policy.ContextWithBufferPolicy(ctx, sessionPolicy.Buffer) + + if newCtx != nil { + ctx = newCtx } 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), - } + var clientSalt [32]byte + clientSaltSlice := clientSalt[:o.method.KeySaltLength] + if _, err := io.ReadFull(rand.Reader, clientSaltSlice); err != nil { + return errors.New("failed to generate client salt").Base(err) } - serverConn := o.method.DialPacketConn(connection) - return singbridge.ReturnError(bufio.CopyPacketConn(ctx, packetConn, serverConn)) + requestDone := func() error { + defer timer.SetTimeout(sessionPolicy.Timeouts.DownlinkOnly) + bufferedWriter := buf.NewBufferedWriter(buf.NewWriter(conn)) + bodyWriter, err := WriteTCPRequest(bufferedWriter, o.method, o.pskList, destination, clientSaltSlice, nil) + if err != nil { + return errors.New("failed to write request").Base(err) + } + + if err = buf.CopyOnceTimeout(link.Reader, bodyWriter, time.Millisecond*100); err != nil && err != buf.ErrNotTimeoutReader && err != buf.ErrReadTimeout { + return errors.New("failed to write A request payload").Base(err).AtWarning() + } + + if err := bufferedWriter.SetBuffered(false); err != nil { + return err + } + + return buf.Copy(link.Reader, bodyWriter, buf.UpdateActivity(timer)) + } + + responseDone := func() error { + defer timer.SetTimeout(sessionPolicy.Timeouts.UplinkOnly) + + responseReader, err := ReadTCPResponse(conn, o.method, o.finalPSK, clientSaltSlice) + if err != nil { + return err + } + + return buf.Copy(responseReader, link.Writer, buf.UpdateActivity(timer)) + } + + responseDoneAndCloseWriter := task.OnSuccess(responseDone, task.Close(link.Writer)) + if err := task.Run(ctx, requestDone, responseDoneAndCloseWriter); err != nil { + return errors.New("connection ends").Base(err) + } + + return nil } + + if network == net.Network_UDP { + requestDone := func() error { + defer timer.SetTimeout(sessionPolicy.Timeouts.DownlinkOnly) + + writer := &UDPWriter{ + Writer: conn, + Destination: destination, + Codec: o.udpCodec, + } + + if err := buf.Copy(link.Reader, writer, buf.UpdateActivity(timer)); err != nil { + return errors.New("failed to transport all UDP request").Base(err) + } + return nil + } + + responseDone := func() error { + defer timer.SetTimeout(sessionPolicy.Timeouts.UplinkOnly) + + reader := &UDPReader{ + Reader: conn, + Codec: o.udpCodec, + } + + if err := buf.Copy(reader, link.Writer, buf.UpdateActivity(timer)); err != nil { + return errors.New("failed to transport all UDP response").Base(err) + } + return nil + } + + responseDoneAndCloseWriter := task.OnSuccess(responseDone, task.Close(link.Writer)) + if err := task.Run(ctx, requestDone, responseDoneAndCloseWriter); err != nil { + return errors.New("connection ends").Base(err) + } + + return nil + } + + return errors.New("unsupported network: ", network) } diff --git a/proxy/shadowsocks_2022/packet.go b/proxy/shadowsocks_2022/packet.go new file mode 100644 index 000000000..d13a9b3a6 --- /dev/null +++ b/proxy/shadowsocks_2022/packet.go @@ -0,0 +1,538 @@ +package shadowsocks_2022 + +import ( + "crypto/cipher" + "crypto/rand" + "encoding/binary" + "io" + "math" + mrand "math/rand/v2" + "sync/atomic" + "time" + + "github.com/xtls/xray-core/common/buf" + "github.com/xtls/xray-core/common/errors" + "github.com/xtls/xray-core/common/net" +) + +type UDPCodec struct { + method *CipherMethod + psk []byte + blockCipher cipher.Block + chachaCipher cipher.AEAD + clientBodyCipher cipher.AEAD + clientSessionID uint64 + nextPacketID atomic.Uint64 + sessions *UDPSessionManager +} + +type UDPPacketCodec = UDPCodec +type UDPServerCodec = UDPCodec + +func newUDPCodec(method *CipherMethod, psk []byte) (*UDPCodec, error) { + c := &UDPCodec{ + method: method, + psk: psk, + } + var err error + if method.IsChaCha { + c.chachaCipher, err = method.NewUDPCipher(psk) + } else { + c.blockCipher, err = method.NewBlock(psk) + } + if err != nil { + return nil, err + } + return c, nil +} + +func NewUDPPacketCodec(method *CipherMethod, psk []byte) (*UDPCodec, error) { + c, err := newUDPCodec(method, psk) + if err != nil { + return nil, err + } + var sessID [8]byte + if _, err := io.ReadFull(rand.Reader, sessID[:]); err != nil { + return nil, err + } + c.clientSessionID = binary.BigEndian.Uint64(sessID[:]) + + if !method.IsChaCha { + clientBodyKey := DeriveSessionSubKey(psk, sessID[:], method.KeySaltLength) + c.clientBodyCipher, err = method.NewAEAD(clientBodyKey) + if err != nil { + return nil, err + } + } + return c, nil +} + +func NewUDPServerCodec(method *CipherMethod, psk []byte, sessionTimeout time.Duration) (*UDPCodec, error) { + c, err := newUDPCodec(method, psk) + if err != nil { + return nil, err + } + c.sessions = NewUDPSessionManager(sessionTimeout) + return c, nil +} + +func (c *UDPCodec) EncodeClientPacket(dest net.Destination, payload []byte) (*buf.Buffer, error) { + packetID := c.nextPacketID.Add(1) + sessID := c.clientSessionID + + // Padding determination (e.g. DNS port 53 disguise) + var paddingLen int + if dest.Port == 53 && len(payload) < MaxPaddingLength { + paddingLen = mrand.IntN(MaxPaddingLength-len(payload)) + 1 + } + + addrPortLen := AddrPortLength(dest) + + if c.method.IsChaCha { + // ChaCha20 mode: 24-byte nonce + plaintext header (27B) + padding + dest + payload + AEAD tag (16B) + totalLen := PacketNonceSize + 27 + paddingLen + addrPortLen + len(payload) + AEADTagSize + if totalLen > buf.Size { + return nil, ErrPacketTooLarge + } + + outBuf := buf.New() + + var nonce [PacketNonceSize]byte + if _, err := io.ReadFull(rand.Reader, nonce[:]); err != nil { + outBuf.Release() + return nil, err + } + outBuf.Write(nonce[:]) + + var hdr [16 + 1 + 8 + 2]byte + binary.BigEndian.PutUint64(hdr[0:8], sessID) + binary.BigEndian.PutUint64(hdr[8:16], packetID) + hdr[16] = HeaderTypeClient + binary.BigEndian.PutUint64(hdr[17:25], uint64(time.Now().Unix())) + binary.BigEndian.PutUint16(hdr[25:27], uint16(paddingLen)) + outBuf.Write(hdr[:]) + if paddingLen > 0 { + outBuf.Write(zeroPadding[:paddingLen]) + } + + if err := WriteAddressPort(outBuf, dest); err != nil { + outBuf.Release() + return nil, err + } + outBuf.Write(payload) + + plainBytes := outBuf.Bytes()[PacketNonceSize:] + outBuf.Extend(int32(c.chachaCipher.Overhead())) + c.chachaCipher.Seal(plainBytes[:0], nonce[:], plainBytes, nil) + return outBuf, nil + } + + // AES mode: + // 16B Encrypted Header + (11B header + padding + dest + payload + 16B AEAD tag) + totalLen := 16 + 11 + paddingLen + addrPortLen + len(payload) + AEADTagSize + if totalLen > buf.Size { + return nil, ErrPacketTooLarge + } + + outBuf := buf.New() + + var rawHeader [16]byte + binary.BigEndian.PutUint64(rawHeader[:8], sessID) + binary.BigEndian.PutUint64(rawHeader[8:16], packetID) + + var encryptedHeader [16]byte + c.blockCipher.Encrypt(encryptedHeader[:], rawHeader[:]) + outBuf.Write(encryptedHeader[:]) + + bodyAead := c.clientBodyCipher + if bodyAead == nil { + bodyKey := DeriveSessionSubKey(c.psk, rawHeader[:8], c.method.KeySaltLength) + var err error + bodyAead, err = c.method.NewAEAD(bodyKey) + if err != nil { + outBuf.Release() + return nil, err + } + } + + var hdr [1 + 8 + 2]byte + hdr[0] = HeaderTypeClient + binary.BigEndian.PutUint64(hdr[1:9], uint64(time.Now().Unix())) + binary.BigEndian.PutUint16(hdr[9:11], uint16(paddingLen)) + outBuf.Write(hdr[:]) + if paddingLen > 0 { + outBuf.Write(zeroPadding[:paddingLen]) + } + + if err := WriteAddressPort(outBuf, dest); err != nil { + outBuf.Release() + return nil, err + } + outBuf.Write(payload) + + plainBytes := outBuf.Bytes()[16:] + bodyNonce := rawHeader[4:16] + outBuf.Extend(int32(bodyAead.Overhead())) + bodyAead.Seal(plainBytes[:0], bodyNonce, plainBytes, nil) + return outBuf, nil +} + +type DecodedUDPPacket struct { + SessionID uint64 + PacketID uint64 + HeaderType byte + Timestamp uint64 + Destination net.Destination + Payload []byte +} + +func parseAddressPort(data []byte) (net.Destination, int, error) { + if len(data) < 1 { + return net.Destination{}, 0, ErrPacketTooShort + } + switch data[0] { + case 1: // IPv4 + if len(data) < 1+4+2 { + return net.Destination{}, 0, ErrPacketTooShort + } + ip := net.IPAddress(data[1:5]) + port := binary.BigEndian.Uint16(data[5:7]) + return net.UDPDestination(ip, net.Port(port)), 7, nil + case 4: // IPv6 + if len(data) < 1+16+2 { + return net.Destination{}, 0, ErrPacketTooShort + } + ip := net.IPAddress(data[1:17]) + port := binary.BigEndian.Uint16(data[17:19]) + return net.UDPDestination(ip, net.Port(port)), 19, nil + case 3: // Domain + if len(data) < 2 { + return net.Destination{}, 0, ErrPacketTooShort + } + domainLen := int(data[1]) + if len(data) < 2+domainLen+2 { + return net.Destination{}, 0, ErrPacketTooShort + } + domain := string(data[2 : 2+domainLen]) + port := binary.BigEndian.Uint16(data[2+domainLen : 2+domainLen+2]) + return net.UDPDestination(net.DomainAddress(domain), net.Port(port)), 2 + domainLen + 2, nil + default: + return net.Destination{}, 0, errors.New("unknown address type") + } +} + +func parsePlainUDPPacket(sessionID, packetID uint64, bodyPlain []byte) (DecodedUDPPacket, error) { + if len(bodyPlain) < 1+8+2 { + return DecodedUDPPacket{}, ErrPacketTooShort + } + + headerType := bodyPlain[0] + epoch := binary.BigEndian.Uint64(bodyPlain[1:9]) + diff := int(math.Abs(float64(time.Now().Unix() - int64(epoch)))) + if diff > 30 { + return DecodedUDPPacket{}, ErrBadTimestamp + } + + offset := 9 + if headerType == HeaderTypeServer { + if len(bodyPlain) < offset+8+2 { + return DecodedUDPPacket{}, ErrPacketTooShort + } + offset += 8 // skip clientSessionID + } + + paddingLen := int(binary.BigEndian.Uint16(bodyPlain[offset : offset+2])) + offset += 2 + + if len(bodyPlain) < offset+paddingLen { + return DecodedUDPPacket{}, ErrNoPadding + } + offset += paddingLen + + dest, addrLen, err := parseAddressPort(bodyPlain[offset:]) + if err != nil { + return DecodedUDPPacket{}, err + } + payload := bodyPlain[offset+addrLen:] + + return DecodedUDPPacket{ + SessionID: sessionID, + PacketID: packetID, + HeaderType: headerType, + Timestamp: epoch, + Destination: dest, + Payload: payload, + }, nil +} + +func (c *UDPCodec) DecodePacket(data []byte) (DecodedUDPPacket, error) { + if len(data) < PacketMinimalHeaderSize { + return DecodedUDPPacket{}, ErrPacketTooShort + } + + if c.method.IsChaCha { + if len(data) < PacketNonceSize+AEADTagSize { + return DecodedUDPPacket{}, ErrPacketTooShort + } + nonce := data[:PacketNonceSize] + ciphertext := data[PacketNonceSize:] + plain, err := c.chachaCipher.Open(ciphertext[:0], nonce, ciphertext, nil) + if err != nil { + return DecodedUDPPacket{}, errors.New("failed to decrypt chacha udp packet").Base(err) + } + if len(plain) < 16+1+8+2 { + return DecodedUDPPacket{}, ErrPacketTooShort + } + + sessionID := binary.BigEndian.Uint64(plain[:8]) + packetID := binary.BigEndian.Uint64(plain[8:16]) + + if c.sessions != nil { + sessionItem, _ := c.sessions.GetOrCreate(sessionID) + sessionItem.Lock() + if !sessionItem.Window.CheckAndAdd(packetID) { + sessionItem.Unlock() + return DecodedUDPPacket{}, ErrPacketIdNotUnique + } + sessionItem.Unlock() + } + + return parsePlainUDPPacket(sessionID, packetID, plain[16:]) + } + + // AES mode + var rawHeader [16]byte + c.blockCipher.Decrypt(rawHeader[:], data[:16]) + sessionID := binary.BigEndian.Uint64(rawHeader[:8]) + packetID := binary.BigEndian.Uint64(rawHeader[8:16]) + + var bodyAead cipher.AEAD + var sessionItem *ServerUDPSession + + if c.sessions != nil { + sessionItem, _ = c.sessions.GetOrCreate(sessionID) + sessionItem.Lock() + if !sessionItem.Window.Check(packetID) { + sessionItem.Unlock() + return DecodedUDPPacket{}, ErrPacketIdNotUnique + } + sessionItem.Unlock() + + bodyAead = sessionItem.GetRemoteCipher() + if bodyAead == nil { + bodyKey := DeriveSessionSubKey(c.psk, rawHeader[:8], c.method.KeySaltLength) + var err error + bodyAead, err = c.method.NewAEAD(bodyKey) + if err != nil { + return DecodedUDPPacket{}, err + } + sessionItem.SetRemoteCipher(bodyAead) + } + } else { + bodyKey := DeriveSessionSubKey(c.psk, rawHeader[:8], c.method.KeySaltLength) + var err error + bodyAead, err = c.method.NewAEAD(bodyKey) + if err != nil { + return DecodedUDPPacket{}, err + } + } + + bodyNonce := rawHeader[4:16] + bodyCipher := data[16:] + bodyPlain, err := bodyAead.Open(bodyCipher[:0], bodyNonce, bodyCipher, nil) + if err != nil { + return DecodedUDPPacket{}, errors.New("failed to decrypt aes udp body").Base(err) + } + + if sessionItem != nil { + sessionItem.Lock() + sessionItem.Window.Add(packetID) + sessionItem.Unlock() + } + + return parsePlainUDPPacket(sessionID, packetID, bodyPlain) +} + +func (c *UDPCodec) Sessions() *UDPSessionManager { + return c.sessions +} + +func (s *ServerUDPSession) EnsureServerState(method *CipherMethod, headerBlock cipher.Block, chachaCipher cipher.AEAD, psk []byte) error { + s.Lock() + defer s.Unlock() + if s.ServerSessionID != 0 { + return nil + } + var sidBuf [8]byte + for { + if _, err := io.ReadFull(rand.Reader, sidBuf[:]); err != nil { + return err + } + s.ServerSessionID = binary.BigEndian.Uint64(sidBuf[:]) + if s.ServerSessionID != 0 { + break + } + } + if method.IsChaCha { + s.ServerChaCha = chachaCipher + } else { + s.ServerBlockCipher = headerBlock + bodyKey := DeriveSessionSubKey(psk, sidBuf[:], method.KeySaltLength) + bodyAead, err := method.NewAEAD(bodyKey) + if err != nil { + s.ServerSessionID = 0 + return err + } + s.ServerCipher = bodyAead + } + return nil +} + +func (s *ServerUDPSession) EncodeServerPacket(method *CipherMethod, clientSessionID uint64, dest net.Destination, payload []byte) ([]byte, error) { + serverSessionID := s.ServerSessionID + serverPacketID := s.ServerPacketID.Add(1) + + if method.IsChaCha { + var nonce [PacketNonceSize]byte + if _, err := io.ReadFull(rand.Reader, nonce[:]); err != nil { + return nil, err + } + + plainBuf := buf.New() + defer plainBuf.Release() + + var hdr [16 + 1 + 8 + 8 + 2]byte + binary.BigEndian.PutUint64(hdr[0:8], serverSessionID) + binary.BigEndian.PutUint64(hdr[8:16], serverPacketID) + hdr[16] = HeaderTypeServer + binary.BigEndian.PutUint64(hdr[17:25], uint64(time.Now().Unix())) + binary.BigEndian.PutUint64(hdr[25:33], clientSessionID) + binary.BigEndian.PutUint16(hdr[33:35], 0) + plainBuf.Write(hdr[:]) + + if err := WriteAddressPort(plainBuf, dest); err != nil { + return nil, err + } + plainBuf.Write(payload) + + sealed := s.ServerChaCha.Seal(nil, nonce[:], plainBuf.Bytes(), nil) + res := make([]byte, PacketNonceSize+len(sealed)) + copy(res[:PacketNonceSize], nonce[:]) + copy(res[PacketNonceSize:], sealed) + return res, nil + } + + // AES mode + var rawHeader [16]byte + binary.BigEndian.PutUint64(rawHeader[:8], serverSessionID) + binary.BigEndian.PutUint64(rawHeader[8:16], serverPacketID) + + var encryptedHeader [16]byte + s.ServerBlockCipher.Encrypt(encryptedHeader[:], rawHeader[:]) + + bodyBuf := buf.New() + defer bodyBuf.Release() + + var hdr [1 + 8 + 8 + 2]byte + hdr[0] = HeaderTypeServer + binary.BigEndian.PutUint64(hdr[1:9], uint64(time.Now().Unix())) + binary.BigEndian.PutUint64(hdr[9:17], clientSessionID) + binary.BigEndian.PutUint16(hdr[17:19], 0) + bodyBuf.Write(hdr[:]) + + if err := WriteAddressPort(bodyBuf, dest); err != nil { + return nil, err + } + bodyBuf.Write(payload) + + bodyNonce := rawHeader[4:16] + sealedBody := s.ServerCipher.Seal(nil, bodyNonce, bodyBuf.Bytes(), nil) + + res := make([]byte, 16+len(sealedBody)) + copy(res[:16], encryptedHeader[:]) + copy(res[16:], sealedBody) + return res, nil +} + +func EncodeServerPacket(method *CipherMethod, headerBlock cipher.Block, chachaAEAD cipher.AEAD, psk []byte, clientSessionID uint64, dest net.Destination, payload []byte) ([]byte, error) { + tempSession := &ServerUDPSession{SessionID: clientSessionID} + if err := tempSession.EnsureServerState(method, headerBlock, chachaAEAD, psk); err != nil { + return nil, err + } + return tempSession.EncodeServerPacket(method, clientSessionID, dest, payload) +} + +func (c *UDPCodec) EncodeServerPacket(clientSessionID uint64, dest net.Destination, payload []byte) ([]byte, error) { + if c.sessions != nil { + sessionItem, _ := c.sessions.GetOrCreate(clientSessionID) + if err := sessionItem.EnsureServerState(c.method, c.blockCipher, c.chachaCipher, c.psk); err != nil { + return nil, err + } + return sessionItem.EncodeServerPacket(c.method, clientSessionID, dest, payload) + } + return EncodeServerPacket(c.method, c.blockCipher, c.chachaCipher, c.psk, clientSessionID, dest, payload) +} + +func (c *UDPCodec) EncodePacket(clientSessionID uint64, dest net.Destination, payload []byte) ([]byte, error) { + return c.EncodeServerPacket(clientSessionID, dest, payload) +} + +type UDPWriter struct { + Writer io.Writer + Destination net.Destination + Codec *UDPPacketCodec +} + +func (w *UDPWriter) WriteMultiBuffer(mb buf.MultiBuffer) error { + for { + mb2, b := buf.SplitFirst(mb) + mb = mb2 + if b == nil { + break + } + dest := w.Destination + if b.UDP != nil { + dest = *b.UDP + } + pktBuf, err := w.Codec.EncodeClientPacket(dest, b.Bytes()) + b.Release() + if err != nil { + buf.ReleaseMulti(mb) + return err + } + _, writeErr := w.Writer.Write(pktBuf.Bytes()) + pktBuf.Release() + if writeErr != nil { + buf.ReleaseMulti(mb) + return writeErr + } + } + return nil +} + +type UDPReader struct { + Reader io.Reader + Codec *UDPPacketCodec +} + +func (r *UDPReader) ReadMultiBuffer() (buf.MultiBuffer, error) { + for { + buffer := buf.New() + _, err := buffer.ReadFrom(r.Reader) + if err != nil { + buffer.Release() + return nil, err + } + + decoded, err := r.Codec.DecodePacket(buffer.Bytes()) + if err != nil { + buffer.Release() + continue + } + buffer.Clear() + buffer.Write(decoded.Payload) + dest := decoded.Destination + buffer.UDP = &dest + return buf.MultiBuffer{buffer}, nil + } +} diff --git a/proxy/shadowsocks_2022/relay_test.go b/proxy/shadowsocks_2022/relay_test.go new file mode 100644 index 000000000..bb8a3277e --- /dev/null +++ b/proxy/shadowsocks_2022/relay_test.go @@ -0,0 +1,230 @@ +package shadowsocks_2022_test + +import ( + "context" + "encoding/base64" + "encoding/binary" + gonet "net" + "sync" + "sync/atomic" + "testing" + "time" + + "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/proxy/shadowsocks_2022" + "github.com/xtls/xray-core/transport" + "lukechampine.com/blake3" +) + +// encodeRelayClientUDPPacket encodes a Shadowsocks-2022 UDP packet with 1 layer of EIH (Relay) +func encodeRelayClientUDPPacket(relayKey, destKey []byte, sessionID, packetID uint64, dest net.Destination, payload []byte) ([]byte, error) { + method, err := GetCipherMethod(MethodAES128GCM) + if err != nil { + return nil, err + } + relayBlock, err := method.NewBlock(relayKey) + if err != nil { + return nil, err + } + + // 1. Plain packet header: sessionID (8B) + packetID (8B) + var rawHeader [16]byte + binary.BigEndian.PutUint64(rawHeader[:8], sessionID) + binary.BigEndian.PutUint64(rawHeader[8:16], packetID) + + // Encrypt packetHeader under relayKey + var encPacketHeader [16]byte + relayBlock.Encrypt(encPacketHeader[:], rawHeader[:]) + + // 2. EI Header: blake3(destKey)[:16] ^ rawHeader + var destHash [16]byte + hash512 := blake3.Sum512(destKey) + copy(destHash[:], hash512[:16]) + + var eiHeader [16]byte + for i := 0; i < 16; i++ { + eiHeader[i] = destHash[i] ^ rawHeader[i] + } + var encEIHeader [16]byte + relayBlock.Encrypt(encEIHeader[:], eiHeader[:]) + + // 3. Payload under destination server's AEAD + bodyKey := DeriveSessionSubKey(destKey, rawHeader[:8], 16) + bodyAead, err := method.NewAEAD(bodyKey) + if err != nil { + return nil, err + } + bodyNonce := rawHeader[4:16] + + outBuf := buf.New() + defer outBuf.Release() + + // VarHeader: client type (1) + timestamp (8) + paddingLen (2) + padding + dest + payload + var hdr [1 + 8 + 2]byte + hdr[0] = HeaderTypeClient + binary.BigEndian.PutUint64(hdr[1:9], uint64(time.Now().Unix())) + binary.BigEndian.PutUint16(hdr[9:11], 0) + outBuf.Write(hdr[:]) + + if err := WriteAddressPort(outBuf, dest); err != nil { + return nil, err + } + outBuf.Write(payload) + + plainBytes := outBuf.Bytes() + outBuf.Extend(int32(bodyAead.Overhead())) + bodyAead.Seal(plainBytes[:0], bodyNonce, plainBytes, nil) + + // Full packet: encPacketHeader (16B) + encEIHeader (16B) + sealedBody + packet := make([]byte, 0, 32+outBuf.Len()) + packet = append(packet, encPacketHeader[:]...) + packet = append(packet, encEIHeader[:]...) + packet = append(packet, outBuf.Bytes()...) + return packet, nil +} + +func TestRelayUDPSessionStabilityAndDispatch(t *testing.T) { + relayKey := []byte("0123456789abcdef") + destKey := []byte("fedcba9876543210") + relayKeyB64 := base64.StdEncoding.EncodeToString(relayKey) + destKeyB64 := base64.StdEncoding.EncodeToString(destKey) + + config := &RelayServerConfig{ + Method: MethodAES128GCM, + Key: relayKeyB64, + Destinations: []*RelayDestination{ + { + Key: destKeyB64, + Address: &net.IPOrDomain{Address: &net.IPOrDomain_Ip{Ip: []byte{127, 0, 0, 1}}}, + Port: 8388, + Email: "dest@example.com", + }, + }, + } + + inbound, err := NewRelayServer(newTestContext(), config) + if err != nil { + t.Fatalf("failed to create RelayServer: %v", err) + } + + sessionID := uint64(0x1122334455667788) + dest := net.UDPDestination(net.LocalHostIP, 8388) + + pkt1, err := encodeRelayClientUDPPacket(relayKey, destKey, sessionID, 1, dest, []byte("xray packet 1")) + if err != nil { + t.Fatalf("failed to encode pkt1: %v", err) + } + pkt2, err := encodeRelayClientUDPPacket(relayKey, destKey, sessionID, 2, dest, []byte("xray packet 2")) + if err != nil { + t.Fatalf("failed to encode pkt2: %v", err) + } + + var dispatchCount atomic.Int32 + var receivedPackets [][]byte + var mu sync.Mutex + + disp := &dummyDispatcher{ + onDispatch: func(ctx context.Context, d net.Destination) (*transport.Link, error) { + dispatchCount.Add(1) + linkR, linkW := gonet.Pipe() + t.Cleanup(func() { + linkW.Close() + linkR.Close() + }) + link := &transport.Link{ + Reader: buf.NewReader(linkR), + Writer: &customWriter{ + write: func(mb buf.MultiBuffer) error { + mu.Lock() + defer mu.Unlock() + for _, b := range mb { + cpy := make([]byte, b.Len()) + copy(cpy, b.Bytes()) + receivedPackets = append(receivedPackets, cpy) + b.Release() + } + return nil + }, + }, + } + return link, nil + }, + } + + clientConn, serverConn := gonet.Pipe() + defer clientConn.Close() + defer serverConn.Close() + + inboundConn := &dummyStatConn{Conn: serverConn} + ctx, cancel := context.WithCancel(newTestContext()) + defer cancel() + + go func() { + _ = inbound.Process(ctx, net.Network_UDP, inboundConn, disp) + }() + + // Send Packet 1 + _, err = clientConn.Write(pkt1) + if err != nil { + t.Fatalf("write pkt1 failed: %v", err) + } + time.Sleep(50 * time.Millisecond) + + // Send Packet 2 (same sessionID, packetID=2) + _, err = clientConn.Write(pkt2) + if err != nil { + t.Fatalf("write pkt2 failed: %v", err) + } + time.Sleep(50 * time.Millisecond) + + // Check dispatch count: For the SAME UDP session, Dispatch MUST be called exactly ONCE! + if count := dispatchCount.Load(); count != 1 { + t.Fatalf("CRITICAL BUG CONFIRMED: expected dispatchCount = 1 for same session, got %d (sessionID was corrupted by Encrypt!)", count) + } + + // Verify downstream destination can decode both packets + method, err := GetCipherMethod(MethodAES128GCM) + common.Must(err) + destCodec, err := NewUDPServerCodec(method, destKey, 300*time.Second) + common.Must(err) + + mu.Lock() + pkts := receivedPackets + mu.Unlock() + + if len(pkts) != 2 { + t.Fatalf("expected 2 received packets at destination, got %d", len(pkts)) + } + + dec1, err := destCodec.DecodePacket(pkts[0]) + if err != nil { + t.Fatalf("dest failed to decode packet 1: %v", err) + } + if dec1.SessionID != sessionID || dec1.PacketID != 1 || string(dec1.Payload) != "xray packet 1" { + t.Fatalf("dec1 mismatch: sess=%x, pktID=%d, payload=%s", dec1.SessionID, dec1.PacketID, string(dec1.Payload)) + } + + dec2, err := destCodec.DecodePacket(pkts[1]) + if err != nil { + t.Fatalf("dest failed to decode packet 2: %v", err) + } + if dec2.SessionID != sessionID || dec2.PacketID != 2 || string(dec2.Payload) != "xray packet 2" { + t.Fatalf("dec2 mismatch: sess=%x, pktID=%d, payload=%s", dec2.SessionID, dec2.PacketID, string(dec2.Payload)) + } +} + +type customWriter struct { + write func(mb buf.MultiBuffer) error +} + +func (w *customWriter) WriteMultiBuffer(mb buf.MultiBuffer) error { + return w.write(mb) +} + +func (w *customWriter) Close() error { + return nil +} + +func (w *customWriter) Interrupt() {} diff --git a/proxy/shadowsocks_2022/replay.go b/proxy/shadowsocks_2022/replay.go new file mode 100644 index 000000000..86a47d8c4 --- /dev/null +++ b/proxy/shadowsocks_2022/replay.go @@ -0,0 +1,159 @@ +package shadowsocks_2022 + +import ( + "crypto/cipher" + "sync" + "sync/atomic" + "time" + + "github.com/xtls/xray-core/common/protocol" + "github.com/xtls/xray-core/common/utils" +) + +const ( + swBlockBitLog = 6 // 1<<6 == 64 bits + swBlockBits = 1 << swBlockBitLog // 64 + swRingBlocks = 1 << 7 // 128 + swBlockMask = swRingBlocks - 1 // 127 + swBitMask = swBlockBits - 1 // 63 + swSize = (swRingBlocks - 1) * swBlockBits // 8128 +) + +type SlidingWindow struct { + last uint64 + ring [swRingBlocks]uint64 +} + +func (f *SlidingWindow) Reset() { + f.last = 0 + f.ring[0] = 0 +} + +func (f *SlidingWindow) Check(counter uint64) bool { + switch { + case counter > f.last: + return true + case f.last-counter > swSize: + return false + } + + blockIndex := (counter >> swBlockBitLog) & swBlockMask + bitIndex := counter & swBitMask + return (f.ring[blockIndex]>>bitIndex)&1 == 0 +} + +func (f *SlidingWindow) Add(counter uint64) { + blockIndex := counter >> swBlockBitLog + + if counter > f.last { + lastBlockIndex := f.last >> swBlockBitLog + diff := int(blockIndex - lastBlockIndex) + if diff > swRingBlocks { + diff = swRingBlocks + } + + for i := 0; i < diff; i++ { + lastBlockIndex = (lastBlockIndex + 1) & swBlockMask + f.ring[lastBlockIndex] = 0 + } + + f.last = counter + } + + blockIndex &= swBlockMask + bitIndex := counter & swBitMask + f.ring[blockIndex] |= 1 << bitIndex +} + +func (f *SlidingWindow) CheckAndAdd(counter uint64) bool { + if !f.Check(counter) { + return false + } + f.Add(counter) + return true +} + +type ServerUDPSession struct { + sync.Mutex + SessionID uint64 + RemoteCipher atomic.Pointer[cipher.AEAD] + Window SlidingWindow + User *protocol.MemoryUser + UserPSK []byte + LastActive atomic.Int64 // Unix timestamp in seconds + + ServerSessionID uint64 + ServerPacketID atomic.Uint64 + ServerCipher cipher.AEAD + ServerBlockCipher cipher.Block + ServerChaCha cipher.AEAD +} + +func (s *ServerUDPSession) GetRemoteCipher() cipher.AEAD { + ptr := s.RemoteCipher.Load() + if ptr == nil { + return nil + } + return *ptr +} + +func (s *ServerUDPSession) SetRemoteCipher(c cipher.AEAD) { + s.RemoteCipher.Store(&c) +} + +type UDPSessionManager struct { + sessions *utils.TypedSyncMap[uint64, *ServerUDPSession] + timeout time.Duration + lastClean atomic.Int64 // Unix timestamp in seconds +} + +func NewUDPSessionManager(timeout time.Duration) *UDPSessionManager { + return &UDPSessionManager{ + sessions: utils.NewTypedSyncMap[uint64, *ServerUDPSession](), + timeout: timeout, + } +} + +func (m *UDPSessionManager) GetOrCreate(sessionID uint64) (*ServerUDPSession, bool) { + now := time.Now().Unix() + if s, ok := m.sessions.Load(sessionID); ok { + s.LastActive.Store(now) + return s, true + } + + s := &ServerUDPSession{ + SessionID: sessionID, + } + s.LastActive.Store(now) + + actual, loaded := m.sessions.LoadOrStore(sessionID, s) + if loaded { + actual.LastActive.Store(now) + return actual, true + } + + // Trigger cleanup if at least 30 seconds have passed since last cleanup + last := m.lastClean.Load() + if now-last > 30 && m.lastClean.CompareAndSwap(last, now) { + go m.cleanup(now) + } + + return s, false +} + +func (m *UDPSessionManager) cleanup(now int64) { + timeoutSec := int64(m.timeout.Seconds()) + if timeoutSec <= 0 { + timeoutSec = 60 + } + m.sessions.Range(func(k uint64, v *ServerUDPSession) bool { + if now-v.LastActive.Load() > timeoutSec { + m.sessions.Delete(k) + } + return true + }) +} + +func (m *UDPSessionManager) Delete(sessionID uint64) { + m.sessions.Delete(sessionID) +} diff --git a/proxy/shadowsocks_2022/shadowsocks_2022.go b/proxy/shadowsocks_2022/shadowsocks_2022.go index 96f62c74b..697cf5955 100644 --- a/proxy/shadowsocks_2022/shadowsocks_2022.go +++ b/proxy/shadowsocks_2022/shadowsocks_2022.go @@ -1 +1,66 @@ package shadowsocks_2022 + +import ( + "context" + "strings" + "sync" + + "github.com/xtls/xray-core/common/errors" + "github.com/xtls/xray-core/common/signal" + "github.com/xtls/xray-core/transport" +) + +type udpConnEntry struct { + sync.Mutex + link *transport.Link + timer *signal.ActivityTimer + cancel context.CancelFunc +} + +const ( + HeaderTypeClient = 0 + HeaderTypeServer = 1 + MaxPaddingLength = 900 + PacketNonceSize = 24 + MaxPacketSize = 65535 + RequestHeaderFixedChunkLength = 1 + 8 + 2 // Type (1B) + Timestamp (8B) + VarHeaderLen (2B) + PacketMinimalHeaderSize = 30 + StreamNonceSize = 12 + AESBlockSize = 16 + AEADTagSize = 16 +) + +var zeroPadding [MaxPaddingLength]byte + +const ( + MethodAES128GCM = "2022-blake3-aes-128-gcm" + MethodAES256GCM = "2022-blake3-aes-256-gcm" + MethodChaCha20Poly1305 = "2022-blake3-chacha20-poly1305" +) + +var List = []string{ + MethodAES128GCM, + MethodAES256GCM, + MethodChaCha20Poly1305, +} + +var ( + ErrBadKey = errors.New("bad key") + ErrBadHeaderType = errors.New("bad header type") + ErrBadTimestamp = errors.New("bad timestamp") + ErrSaltNotUnique = errors.New("salt not unique") + ErrPacketIdNotUnique = errors.New("packet id not unique") + ErrPacketTooShort = errors.New("packet too short") + ErrPacketTooLarge = errors.New("packet too large") + ErrNoPadding = errors.New("bad request: missing payload or padding") + ErrInvalidRequest = errors.New("invalid request") +) + +func IsSupportedMethod(method string) bool { + for _, m := range List { + if strings.EqualFold(m, method) { + return true + } + } + return false +} diff --git a/proxy/shadowsocks_2022/shadowsocks_2022_test.go b/proxy/shadowsocks_2022/shadowsocks_2022_test.go new file mode 100644 index 000000000..0f8df2c49 --- /dev/null +++ b/proxy/shadowsocks_2022/shadowsocks_2022_test.go @@ -0,0 +1,794 @@ +package shadowsocks_2022_test + +import ( + "bytes" + "context" + "crypto/rand" + "encoding/base64" + "encoding/binary" + "io" + gonet "net" + "sync" + "testing" + "time" + + "github.com/google/go-cmp/cmp" + "github.com/xtls/xray-core/common" + "github.com/xtls/xray-core/common/antireplay" + "github.com/xtls/xray-core/common/buf" + "github.com/xtls/xray-core/common/errors" + "github.com/xtls/xray-core/common/net" + "github.com/xtls/xray-core/common/protocol" + "github.com/xtls/xray-core/common/serial" + "github.com/xtls/xray-core/common/session" + "github.com/xtls/xray-core/core" + "github.com/xtls/xray-core/features/routing" + . "github.com/xtls/xray-core/proxy/shadowsocks_2022" + "github.com/xtls/xray-core/transport" + "github.com/xtls/xray-core/transport/internet/stat" +) + +func newTestContext() context.Context { + v, err := core.New(&core.Config{}) + common.Must(err) + ctx := context.WithValue(context.Background(), core.XrayKey(1), v) + ctx = session.ContextWithInbound(ctx, &session.Inbound{}) + return ctx +} + +func generateRandomKey(size int) string { + b := make([]byte, size) + _, _ = rand.Read(b) + return base64.StdEncoding.EncodeToString(b) +} + +func TestKDF(t *testing.T) { + // Test ParseKey + if _, err := ParseKey("", 16); err != ErrBadKey { + t.Fatalf("expected ErrBadKey for empty key, got %v", err) + } + + shortKey := base64.StdEncoding.EncodeToString([]byte("short")) + if _, err := ParseKey(shortKey, 16); err != ErrBadKey { + t.Fatalf("expected ErrBadKey for short key, got %v", err) + } + + exactKey := []byte("0123456789abcdef") + exactKeyB64 := base64.StdEncoding.EncodeToString(exactKey) + normExact, err := ParseKey(exactKeyB64, 16) + if err != nil || !bytes.Equal(normExact, exactKey) { + t.Fatalf("unexpected parsed exact key: %v, err: %v", normExact, err) + } + + longKey := base64.StdEncoding.EncodeToString([]byte("0123456789abcdef_longer_key_for_testing")) + if _, err := ParseKey(longKey, 16); err != ErrBadKey { + t.Fatalf("expected ErrBadKey for long key, got %v", err) + } + + // Test Session Subkey determinism + salt := []byte("random_salt_1234") + k1 := DeriveSessionSubKey(normExact, salt, 16) + k2 := DeriveSessionSubKey(normExact, salt, 16) + if !bytes.Equal(k1, k2) { + t.Fatal("DeriveSessionSubKey should be deterministic") + } + + // Identity subkey must differ from session subkey with same inputs + idKey := DeriveIdentitySubKey(normExact, salt, 16) + if bytes.Equal(k1, idKey) { + t.Fatal("DeriveIdentitySubKey must differ from DeriveSessionSubKey") + } + + // User PSK hash + h1 := DeriveUserPSKHash(normExact) + h2 := DeriveUserPSKHash(normExact) + if h1 != h2 { + t.Fatal("DeriveUserPSKHash should be deterministic") + } +} + +func TestReplayFilter(t *testing.T) { + filter := antireplay.NewMapFilter[string](60) + + salt1 := []byte("test_salt_111111") + salt2 := []byte("test_salt_222222") + + if !filter.Check(string(salt1)) { + t.Fatal("first check on salt1 should be true") + } + if filter.Check(string(salt1)) { + t.Fatal("second check on salt1 should be false (replay detected)") + } + + if !filter.Check(string(salt2)) { + t.Fatal("first check on salt2 should be true") + } + + // Test SlidingWindow + var window SlidingWindow + if !window.Check(1) { + t.Fatal("packet 1 should be accepted") + } + window.Add(1) + + if window.Check(1) { + t.Fatal("duplicate packet 1 should be rejected") + } + + if !window.Check(100) { + t.Fatal("packet 100 should be accepted") + } + window.Add(100) + + if window.Check(100) { + t.Fatal("duplicate packet 100 should be rejected") + } + + if !window.Check(50) { + t.Fatal("out-of-order packet 50 within window should be accepted") + } + window.Add(50) + if window.Check(50) { + t.Fatal("duplicate packet 50 should be rejected") + } + + // Check packet far behind window (> 8128) + window.Add(10000) + if window.Check(1) { + t.Fatal("packet 1 should be rejected as behind window") + } +} + +func TestTCPStreamAndHandshake(t *testing.T) { + methods := []struct { + name string + keySize int + }{ + {MethodAES128GCM, 16}, + {MethodAES256GCM, 32}, + {MethodChaCha20Poly1305, 32}, + } + + dest := net.TCPDestination(net.LocalHostIP, net.Port(8080)) + testPayload := []byte("Hello, Shadowsocks 2022 Native Implementation!") + + for _, m := range methods { + t.Run(m.name, func(t *testing.T) { + rawKey := make([]byte, m.keySize) + _, _ = rand.Read(rawKey) + method, err := GetCipherMethod(m.name) + common.Must(err) + + clientConn, serverConn := gonet.Pipe() + defer clientConn.Close() + defer serverConn.Close() + + var wg sync.WaitGroup + wg.Add(2) + + var receivedDest net.Destination + var receivedPayload []byte + + // Server goroutine + go func() { + defer wg.Done() + salt := make([]byte, method.KeySaltLength) + _, err := io.ReadFull(serverConn, salt) + common.Must(err) + + sessionKey := DeriveSessionSubKey(rawKey, salt, method.KeySaltLength) + aead, err := method.NewAEAD(sessionKey) + common.Must(err) + + reader := NewStreamReader(serverConn, aead) + + // Read fixed chunk (11 + 16 bytes) + var fixedBuf [RequestHeaderFixedChunkLength + AEADTagSize]byte + _, err = io.ReadFull(serverConn, fixedBuf[:]) + common.Must(err) + + plainFixed, err := aead.Open(fixedBuf[:0], reader.Nonce(), fixedBuf[:], nil) + common.Must(err) + IncreaseNonce(reader.Nonce()) + if plainFixed[0] != HeaderTypeClient { + t.Errorf("expected client header type, got %d", plainFixed[0]) + } + + // Read variable chunk + varLen := int(plainFixed[9])<<8 | int(plainFixed[10]) + varBuf := make([]byte, varLen+AEADTagSize) + _, err = io.ReadFull(serverConn, varBuf) + common.Must(err) + + plainVar, err := aead.Open(varBuf[:0], reader.Nonce(), varBuf, nil) + common.Must(err) + IncreaseNonce(reader.Nonce()) + + vBuf := buf.New() + vBuf.Write(plainVar) + receivedDest, err = ReadAddressPort(vBuf) + common.Must(err) + + // Skip padding + var padBytes [2]byte + _, _ = vBuf.Read(padBytes[:]) + padLen := int(padBytes[0])<<8 | int(padBytes[1]) + vBuf.Advance(int32(padLen)) + + receivedPayload = make([]byte, vBuf.Len()) + copy(receivedPayload, vBuf.Bytes()) + vBuf.Release() + + // Server sends response handshake + serverSalt := make([]byte, method.KeySaltLength) + _, _ = rand.Read(serverSalt) + respKey := DeriveSessionSubKey(rawKey, serverSalt, method.KeySaltLength) + respAead, err := method.NewAEAD(respKey) + writer := NewStreamWriter(serverConn, respAead) + _, _ = serverConn.Write(serverSalt) + + fixedResp := make([]byte, 1+8+method.KeySaltLength+2) + fixedResp[0] = HeaderTypeServer + binary.BigEndian.PutUint64(fixedResp[1:9], uint64(time.Now().Unix())) + copy(fixedResp[9:9+method.KeySaltLength], salt) + binary.BigEndian.PutUint16(fixedResp[9+method.KeySaltLength:11+method.KeySaltLength], 0) + + fixedChunk := respAead.Seal(nil, writer.Nonce(), fixedResp, nil) + IncreaseNonce(writer.Nonce()) + _, _ = serverConn.Write(fixedChunk) + + // Echo stream data + mb, err := reader.ReadMultiBuffer() + common.Must(err) + _ = writer.WriteMultiBuffer(mb) + }() + + // Client goroutine + go func() { + defer wg.Done() + clientSalt, writer, err := ClientHandshake(clientConn, method, [][]byte{rawKey}, dest, testPayload) + common.Must(err) + + reader, _, err := ClientVerifyServerResponse(clientConn, method, rawKey, clientSalt) + common.Must(err) + + // Send additional stream data + streamData := []byte("stream chunk test") + _ = writer.WriteChunk(streamData) + + mb, err := reader.ReadMultiBuffer() + common.Must(err) + if !bytes.Equal(mb[0].Bytes(), streamData) { + t.Errorf("echoed stream data mismatch: got %s, want %s", mb[0].Bytes(), streamData) + } + buf.ReleaseMulti(mb) + }() + + wg.Wait() + + if receivedDest.NetAddr() != dest.NetAddr() { + t.Errorf("destination mismatch: got %s, want %s", receivedDest.NetAddr(), dest.NetAddr()) + } + if diff := cmp.Diff(receivedPayload, testPayload); diff != "" { + t.Errorf("payload mismatch: %s", diff) + } + }) + } +} + +func TestUDPCodec(t *testing.T) { + methods := []string{ + MethodAES128GCM, + MethodAES256GCM, + MethodChaCha20Poly1305, + } + + dest := net.UDPDestination(net.LocalHostIP, net.Port(53)) + payload := []byte("DNS query payload") + + for _, methodName := range methods { + t.Run(methodName, func(t *testing.T) { + method, err := GetCipherMethod(methodName) + common.Must(err) + + psk := make([]byte, method.KeySaltLength) + _, _ = rand.Read(psk) + + clientCodec, err := NewUDPPacketCodec(method, psk) + common.Must(err) + serverCodec, err := NewUDPServerCodec(method, psk, time.Minute) + common.Must(err) + + pktBuf, err := clientCodec.EncodeClientPacket(dest, payload) + common.Must(err) + defer pktBuf.Release() + + decoded, err := serverCodec.DecodePacket(pktBuf.Bytes()) + common.Must(err) + + if decoded.HeaderType != HeaderTypeClient { + t.Errorf("expected header type %d, got %d", HeaderTypeClient, decoded.HeaderType) + } + if decoded.Destination.Port != dest.Port { + t.Errorf("port mismatch: got %d, want %d", decoded.Destination.Port, dest.Port) + } + if !bytes.Equal(decoded.Payload, payload) { + t.Errorf("payload mismatch: got %s, want %s", decoded.Payload, payload) + } + }) + } +} + +func TestMultiUserManager(t *testing.T) { + masterKey := generateRandomKey(16) + userKey1 := generateRandomKey(16) + userKey2 := generateRandomKey(16) + + config := &MultiUserServerConfig{ + Method: MethodAES128GCM, + Key: masterKey, + Users: []*protocol.User{ + { + Email: "user1@example.com", + Account: serial.ToTypedMessage(&Account{Key: userKey1}), + }, + }, + } + + inbound, err := NewMultiServer(newTestContext(), config) + common.Must(err) + + if inbound.GetUsersCount(context.Background()) != 1 { + t.Fatalf("expected 1 user, got %d", inbound.GetUsersCount(context.Background())) + } + + u1 := inbound.GetUser(context.Background(), "user1@example.com") + if u1 == nil || u1.Email != "user1@example.com" { + t.Fatal("user1 not found") + } + + // Add User 2 + rawKey2, _ := base64.StdEncoding.DecodeString(userKey2) + u2 := &protocol.MemoryUser{ + Email: "user2@example.com", + Account: &MemoryAccount{ + Key: rawKey2, + }, + } + err = inbound.AddUser(context.Background(), u2) + common.Must(err) + + if inbound.GetUsersCount(context.Background()) != 2 { + t.Fatalf("expected 2 users, got %d", inbound.GetUsersCount(context.Background())) + } + + // Remove User 1 + err = inbound.RemoveUser(context.Background(), "user1@example.com") + common.Must(err) + + if inbound.GetUsersCount(context.Background()) != 1 { + t.Fatalf("expected 1 user, got %d", inbound.GetUsersCount(context.Background())) + } + if inbound.GetUser(context.Background(), "user1@example.com") != nil { + t.Fatal("user1 should have been removed") + } +} + +type dummyDispatcher struct { + onDispatch func(ctx context.Context, dest net.Destination) (*transport.Link, error) +} + +func (d *dummyDispatcher) Dispatch(ctx context.Context, dest net.Destination) (*transport.Link, error) { + if d.onDispatch != nil { + return d.onDispatch(ctx, dest) + } + return nil, errors.New("not handled") +} + +func (d *dummyDispatcher) DispatchLink(ctx context.Context, dest net.Destination, link *transport.Link) error { + return nil +} + +func (d *dummyDispatcher) Start() error { return nil } +func (d *dummyDispatcher) Close() error { return nil } +func (d *dummyDispatcher) Type() interface{} { return routing.DispatcherType() } + +type dummyStatConn struct { + gonet.Conn +} + +func (c *dummyStatConn) ReadMultiBuffer() (buf.MultiBuffer, error) { + b := buf.New() + _, err := b.ReadFrom(c.Conn) + return buf.MultiBuffer{b}, err +} + +func (c *dummyStatConn) WriteMultiBuffer(mb buf.MultiBuffer) error { + defer buf.ReleaseMulti(mb) + for _, b := range mb { + if _, err := c.Conn.Write(b.Bytes()); err != nil { + return err + } + } + return nil +} + +func TestMultiUserTCPConnection(t *testing.T) { + masterKey := generateRandomKey(16) + userKey1 := generateRandomKey(16) + userKey2 := generateRandomKey(16) + + config := &MultiUserServerConfig{ + Method: MethodAES128GCM, + Key: masterKey, + Users: []*protocol.User{ + { + Email: "user1@example.com", + Account: serial.ToTypedMessage(&Account{Key: userKey1}), + }, + { + Email: "user2@example.com", + Account: serial.ToTypedMessage(&Account{Key: userKey2}), + }, + }, + } + + testCtx := newTestContext() + inbound, err := NewMultiServer(testCtx, config) + common.Must(err) + + clientConn, serverConn := gonet.Pipe() + defer clientConn.Close() + defer serverConn.Close() + + dest := net.TCPDestination(net.LocalHostIP, 443) + method, err := GetCipherMethod(MethodAES128GCM) + common.Must(err) + + masterRaw, _ := base64.StdEncoding.DecodeString(masterKey) + user2Raw, _ := base64.StdEncoding.DecodeString(userKey2) + clientPSKList := [][]byte{masterRaw, user2Raw} + + dispatchedUserChan := make(chan string, 1) + + disp := &dummyDispatcher{ + onDispatch: func(ctx context.Context, d net.Destination) (*transport.Link, error) { + inbound := session.InboundFromContext(ctx) + if inbound != nil && inbound.User != nil { + dispatchedUserChan <- inbound.User.Email + } + link := &transport.Link{ + Reader: buf.NewReader(bytes.NewReader(nil)), + Writer: buf.Discard, + } + return link, nil + }, + } + + go func() { + _ = inbound.Process(testCtx, net.Network_TCP, &dummyStatConn{Conn: serverConn}, disp) + }() + + clientSalt, writer, err := ClientHandshake(clientConn, method, clientPSKList, dest, []byte("ping")) + common.Must(err) + + reader, _, err := ClientVerifyServerResponse(clientConn, method, user2Raw, clientSalt) + common.Must(err) + _ = writer + _ = reader + + select { + case email := <-dispatchedUserChan: + if email != "user2@example.com" { + t.Fatalf("expected user2@example.com, got %s", email) + } + case <-time.After(2 * time.Second): + t.Fatal("timeout waiting for dispatched user") + } +} + +func TestUDPReaderWriter(t *testing.T) { + for _, methodName := range []string{MethodAES128GCM, MethodAES256GCM, MethodChaCha20Poly1305} { + t.Run(methodName, func(t *testing.T) { + method, err := GetCipherMethod(methodName) + common.Must(err) + rawPSK := make([]byte, method.KeySaltLength) + _, _ = rand.Read(rawPSK) + + clientCodec, err := NewUDPPacketCodec(method, rawPSK) + common.Must(err) + serverCodec, err := NewUDPServerCodec(method, rawPSK, time.Minute) + common.Must(err) + + dest := net.UDPDestination(net.LocalHostIP, 53) + + // Client to Server + clientPacketBuf, err := clientCodec.EncodeClientPacket(dest, []byte("hello dns")) + common.Must(err) + defer clientPacketBuf.Release() + + serverDecoded, err := serverCodec.DecodePacket(clientPacketBuf.Bytes()) + common.Must(err) + if string(serverDecoded.Payload) != "hello dns" { + t.Fatalf("unexpected server decoded payload: %s", string(serverDecoded.Payload)) + } + + // Server to Client + serverPacket, err := serverCodec.EncodePacket(serverDecoded.SessionID, dest, []byte("dns response")) + common.Must(err) + + clientDecoded, err := clientCodec.DecodePacket(serverPacket) + common.Must(err) + if string(clientDecoded.Payload) != "dns response" { + t.Fatalf("unexpected client decoded payload: %s", string(clientDecoded.Payload)) + } + + // Test UDPWriter and UDPReader pipeline + pipeR, pipeW := gonet.Pipe() + defer pipeR.Close() + defer pipeW.Close() + + writer := &UDPWriter{ + Writer: pipeW, + Destination: dest, + Codec: clientCodec, + } + reader := &UDPReader{ + Reader: pipeR, + Codec: clientCodec, + } + + go func() { + // Simulate server echoing back as server response + buf := make([]byte, 2048) + n, err := pipeR.Read(buf) + if err != nil { + return + } + dec, err := serverCodec.DecodePacket(buf[:n]) + if err != nil { + return + } + resp, err := serverCodec.EncodePacket(dec.SessionID, dest, dec.Payload) + if err != nil { + return + } + _, _ = pipeW.Write(resp) + }() + + b := buf.New() + b.WriteString("piped udp packet") + common.Must(writer.WriteMultiBuffer(buf.MultiBuffer{b})) + + received, err := reader.ReadMultiBuffer() + common.Must(err) + if received[0].String() != "piped udp packet" { + t.Fatalf("expected 'piped udp packet', got '%s'", received[0].String()) + } + }) + } +} + +func TestTCPRequestResponse(t *testing.T) { + for _, methodName := range []string{MethodAES128GCM, MethodAES256GCM, MethodChaCha20Poly1305} { + t.Run(methodName, func(t *testing.T) { + method, err := GetCipherMethod(methodName) + common.Must(err) + rawPSK := make([]byte, method.KeySaltLength) + _, _ = rand.Read(rawPSK) + + clientConn, serverConn := gonet.Pipe() + defer clientConn.Close() + defer serverConn.Close() + + clientSalt := make([]byte, method.KeySaltLength) + _, _ = rand.Read(clientSalt) + dest := net.TCPDestination(net.LocalHostIP, 80) + + go func() { + // Server side: read handshake and verify clientSalt + salt := make([]byte, method.KeySaltLength) + if _, err := io.ReadFull(serverConn, salt); err != nil { + t.Errorf("server read salt error: %v", err) + return + } + sessionKey := DeriveSessionSubKey(rawPSK, salt, method.KeySaltLength) + aead, err := method.NewAEAD(sessionKey) + if err != nil { + t.Errorf("server AEAD error: %v", err) + return + } + sReader := NewStreamReader(serverConn, aead) + var fixedBuf [RequestHeaderFixedChunkLength + AEADTagSize]byte + if _, err := io.ReadFull(serverConn, fixedBuf[:]); err != nil { + t.Errorf("server read fixed error: %v", err) + return + } + plainFixed, err := aead.Open(fixedBuf[:0], sReader.Nonce(), fixedBuf[:], nil) + if err != nil { + t.Errorf("server decrypt fixed error: %v", err) + return + } + IncreaseNonce(sReader.Nonce()) + + varLen := int(binary.BigEndian.Uint16(plainFixed[9:11])) + varBuf := make([]byte, varLen+AEADTagSize) + if _, err := io.ReadFull(serverConn, varBuf); err != nil { + t.Errorf("server read var error: %v", err) + return + } + plainVar, err := aead.Open(varBuf[:0], sReader.Nonce(), varBuf, nil) + if err != nil { + t.Errorf("server decrypt var error: %v", err) + return + } + IncreaseNonce(sReader.Nonce()) + + vBuf := buf.New() + vBuf.Write(plainVar) + receivedDest, err := ReadAddressPort(vBuf) + if err != nil || receivedDest != dest { + t.Errorf("dest mismatch: %v vs %v, err: %v", receivedDest, dest, err) + return + } + + // Echo client salt back to client using WriteTCPResponse + sWriter, err := WriteTCPResponse(serverConn, method, rawPSK, salt, []byte("early-reply")) + if err != nil { + t.Errorf("server response error: %v", err) + return + } + _ = sWriter + }() + + bodyWriter, err := WriteTCPRequest(clientConn, method, [][]byte{rawPSK}, dest, clientSalt, nil) + common.Must(err) + _ = bodyWriter + + responseReader, err := ReadTCPResponse(clientConn, method, rawPSK, clientSalt) + common.Must(err) + + mb, err := responseReader.ReadMultiBuffer() + common.Must(err) + if mb[0].String() != "early-reply" { + t.Fatalf("expected early-reply, got %s", mb[0].String()) + } + }) + } +} + +func TestUDPReplayProtection(t *testing.T) { + for _, methodName := range []string{MethodAES128GCM, MethodAES256GCM, MethodChaCha20Poly1305} { + t.Run(methodName, func(t *testing.T) { + method, err := GetCipherMethod(methodName) + common.Must(err) + rawPSK := make([]byte, method.KeySaltLength) + _, _ = rand.Read(rawPSK) + + clientCodec, err := NewUDPPacketCodec(method, rawPSK) + common.Must(err) + serverCodec, err := NewUDPServerCodec(method, rawPSK, time.Minute) + common.Must(err) + + dest := net.UDPDestination(net.LocalHostIP, 53) + pktBuf, err := clientCodec.EncodeClientPacket(dest, []byte("dns 1")) + common.Must(err) + defer pktBuf.Release() + + rawCopy := make([]byte, pktBuf.Len()) + copy(rawCopy, pktBuf.Bytes()) + + // First decode should succeed + _, err = serverCodec.DecodePacket(pktBuf.Bytes()) + if err != nil { + t.Fatalf("first decode failed: %v", err) + } + + // Replay same packet wire bytes should fail with ErrPacketIdNotUnique + _, err = serverCodec.DecodePacket(rawCopy) + if err != ErrPacketIdNotUnique { + t.Fatalf("expected ErrPacketIdNotUnique on replay, got: %v", err) + } + }) + } +} + +func TestServerUDPSessionStabilityAndMonotonicPacketID(t *testing.T) { + for _, methodName := range []string{MethodAES128GCM, MethodAES256GCM, MethodChaCha20Poly1305} { + t.Run(methodName, func(t *testing.T) { + method, err := GetCipherMethod(methodName) + common.Must(err) + rawPSK := make([]byte, method.KeySaltLength) + _, _ = rand.Read(rawPSK) + + clientCodec, err := NewUDPPacketCodec(method, rawPSK) + common.Must(err) + serverCodec, err := NewUDPServerCodec(method, rawPSK, time.Minute) + common.Must(err) + + dest := net.UDPDestination(net.LocalHostIP, 53) + + // Client sends packet 1 + pkt1, err := clientCodec.EncodeClientPacket(dest, []byte("request 1")) + common.Must(err) + defer pkt1.Release() + + dec1, err := serverCodec.DecodePacket(pkt1.Bytes()) + common.Must(err) + + // Server sends response 1 + resp1, err := serverCodec.EncodePacket(dec1.SessionID, dest, []byte("response 1")) + common.Must(err) + + // Server sends response 2 to the same client session + resp2, err := serverCodec.EncodePacket(dec1.SessionID, dest, []byte("response 2")) + common.Must(err) + + // Decode both on client + cDec1, err := clientCodec.DecodePacket(resp1) + common.Must(err) + cDec2, err := clientCodec.DecodePacket(resp2) + common.Must(err) + + if cDec1.SessionID != cDec2.SessionID { + t.Fatalf("expected stable server session ID, got %d and %d", cDec1.SessionID, cDec2.SessionID) + } + if cDec2.PacketID <= cDec1.PacketID { + t.Fatalf("expected monotonically increasing packet ID, got %d then %d", cDec1.PacketID, cDec2.PacketID) + } + if string(cDec1.Payload) != "response 1" || string(cDec2.Payload) != "response 2" { + t.Fatalf("payload mismatch") + } + }) + } +} + +type mockDialer struct { + dial func(ctx context.Context, dest net.Destination) (stat.Connection, error) +} + +func (d *mockDialer) Dial(ctx context.Context, dest net.Destination) (stat.Connection, error) { + if d.dial != nil { + return d.dial(ctx, dest) + } + c1, c2 := gonet.Pipe() + _ = c2.Close() + return &dummyStatConn{Conn: c1}, nil +} + +func (d *mockDialer) DestIpAddress() net.IP { + return net.IP{127, 0, 0, 1} +} + +func (d *mockDialer) SetOutboundGateway(ctx context.Context, ob *session.Outbound) {} + +func TestOutboundProcess(t *testing.T) { + testCtx := newTestContext() + key := generateRandomKey(16) + clientConfig := &ClientConfig{ + Address: &net.IPOrDomain{Address: &net.IPOrDomain_Ip{Ip: []byte{127, 0, 0, 1}}}, + Port: 1080, + Method: MethodAES128GCM, + Key: key, + } + + outbound, err := NewClient(testCtx, clientConfig) + common.Must(err) + + ctx := session.ContextWithOutbounds(testCtx, []*session.Outbound{ + { + Target: net.TCPDestination(net.LocalHostIP, 80), + }, + }) + + link := &transport.Link{ + Reader: buf.NewReader(bytes.NewReader(nil)), + Writer: buf.Discard, + } + + dialer := &mockDialer{} + err = outbound.Process(ctx, link, dialer) + if err == nil { + t.Fatal("expected error from closed mock dialer pipe, got nil") + } +} diff --git a/proxy/shadowsocks_2022/stream.go b/proxy/shadowsocks_2022/stream.go new file mode 100644 index 000000000..87a3668c5 --- /dev/null +++ b/proxy/shadowsocks_2022/stream.go @@ -0,0 +1,545 @@ +package shadowsocks_2022 + +import ( + "crypto/cipher" + "crypto/rand" + "encoding/binary" + "io" + "math" + mrand "math/rand" + "time" + + "github.com/xtls/xray-core/common/buf" + "github.com/xtls/xray-core/common/errors" + "github.com/xtls/xray-core/common/net" + "github.com/xtls/xray-core/common/protocol" +) + +var addrParser = protocol.NewAddressParser( + protocol.AddressFamilyByte(0x01, net.AddressFamilyIPv4), + protocol.AddressFamilyByte(0x04, net.AddressFamilyIPv6), + protocol.AddressFamilyByte(0x03, net.AddressFamilyDomain), + protocol.WithAddressTypeParser(func(b byte) byte { + return b & 0x0F + }), +) + +func IncreaseNonce(nonce []byte) { + for i := range nonce { + nonce[i]++ + if nonce[i] != 0 { + return + } + } +} + +// WriteAddressPort writes a destination address and port in SOCKS5 format +func WriteAddressPort(w io.Writer, dest net.Destination) error { + return addrParser.WriteAddressPort(w, dest.Address, dest.Port) +} + +// ReadAddressPort reads a destination address and port in SOCKS5 format +func ReadAddressPort(r io.Reader) (net.Destination, error) { + addr, port, err := addrParser.ReadAddressPort(nil, r) + if err != nil { + return net.Destination{}, err + } + return net.TCPDestination(addr, port), nil +} + +// AddrPortLength returns the serialized length of a destination in SOCKS5 format +func AddrPortLength(dest net.Destination) int { + switch dest.Address.Family() { + case net.AddressFamilyIPv4: + return 1 + 4 + 2 + case net.AddressFamilyDomain: + return 1 + 1 + len(dest.Address.Domain()) + 2 + case net.AddressFamilyIPv6: + return 1 + 16 + 2 + default: + return 0 + } +} + +type StreamWriter struct { + writer io.Writer + cipher cipher.AEAD + nonce [StreamNonceSize]byte + lenBuf [2]byte + buf []byte +} + +func NewStreamWriter(w io.Writer, c cipher.AEAD) *StreamWriter { + return &StreamWriter{ + writer: w, + cipher: c, + buf: make([]byte, 0, MaxPacketSize+2+2*AEADTagSize), + } +} + +func (w *StreamWriter) Nonce() []byte { + return w.nonce[:] +} + +func (w *StreamWriter) Cipher() cipher.AEAD { + return w.cipher +} + +func (w *StreamWriter) WriteChunk(payload []byte) error { + payloadLen := len(payload) + if payloadLen == 0 { + return nil + } + if payloadLen > MaxPacketSize { + return errors.New("payload exceeds MaxPacketSize") + } + + totalSize := 2 + AEADTagSize + payloadLen + AEADTagSize + if cap(w.buf) < totalSize { + w.buf = make([]byte, 0, totalSize) + } + + binary.BigEndian.PutUint16(w.lenBuf[:], uint16(payloadLen)) + w.buf = w.cipher.Seal(w.buf[:0], w.nonce[:], w.lenBuf[:], nil) + IncreaseNonce(w.nonce[:]) + + w.buf = w.cipher.Seal(w.buf, w.nonce[:], payload, nil) + IncreaseNonce(w.nonce[:]) + + _, err := w.writer.Write(w.buf) + return err +} + +func (w *StreamWriter) Write(p []byte) (int, error) { + n := len(p) + for len(p) > 0 { + chunkSize := len(p) + if chunkSize > MaxPacketSize { + chunkSize = MaxPacketSize + } + if err := w.WriteChunk(p[:chunkSize]); err != nil { + return 0, err + } + p = p[chunkSize:] + } + return n, nil +} + +func (w *StreamWriter) WriteMultiBuffer(mb buf.MultiBuffer) error { + defer buf.ReleaseMulti(mb) + for _, b := range mb { + if err := w.WriteChunk(b.Bytes()); err != nil { + return err + } + } + return nil +} + +type StreamReader struct { + reader io.Reader + cipher cipher.AEAD + nonce [StreamNonceSize]byte + lenBuf [2 + AEADTagSize]byte + buffer []byte + cached int + offset int +} + +func NewStreamReader(r io.Reader, c cipher.AEAD) *StreamReader { + return &StreamReader{ + reader: r, + cipher: c, + buffer: make([]byte, MaxPacketSize+AEADTagSize), + } +} + +func (r *StreamReader) Nonce() []byte { + return r.nonce[:] +} + +func (r *StreamReader) Cipher() cipher.AEAD { + return r.cipher +} + +func (r *StreamReader) Read(p []byte) (int, error) { + if r.cached > 0 { + n := copy(p, r.buffer[r.offset:r.offset+r.cached]) + r.cached -= n + r.offset += n + return n, nil + } + + // Read 2-byte length + AEAD tag (18 bytes) + if _, err := io.ReadFull(r.reader, r.lenBuf[:]); err != nil { + return 0, err + } + + decryptedLen, err := r.cipher.Open(r.lenBuf[:0], r.nonce[:], r.lenBuf[:], nil) + if err != nil { + return 0, errors.New("failed to decrypt chunk length").Base(err) + } + IncreaseNonce(r.nonce[:]) + + payloadLen := int(binary.BigEndian.Uint16(decryptedLen)) + if payloadLen == 0 { + return 0, nil + } + + chunkEnd := payloadLen + AEADTagSize + if _, err := io.ReadFull(r.reader, r.buffer[:chunkEnd]); err != nil { + return 0, err + } + + decryptedPayload, err := r.cipher.Open(r.buffer[:0], r.nonce[:], r.buffer[:chunkEnd], nil) + if err != nil { + return 0, errors.New("failed to decrypt chunk payload").Base(err) + } + IncreaseNonce(r.nonce[:]) + + r.cached = len(decryptedPayload) + r.offset = 0 + + n := copy(p, r.buffer[r.offset:r.offset+r.cached]) + r.cached -= n + r.offset += n + return n, nil +} + +func (r *StreamReader) ReadMultiBuffer() (buf.MultiBuffer, error) { + if r.cached > 0 { + b := buf.New() + b.Write(r.buffer[r.offset : r.offset+r.cached]) + r.cached = 0 + r.offset = 0 + return buf.MultiBuffer{b}, nil + } + + if _, err := io.ReadFull(r.reader, r.lenBuf[:]); err != nil { + return nil, err + } + + decryptedLen, err := r.cipher.Open(r.lenBuf[:0], r.nonce[:], r.lenBuf[:], nil) + if err != nil { + return nil, errors.New("failed to decrypt chunk length").Base(err) + } + IncreaseNonce(r.nonce[:]) + + payloadLen := int(binary.BigEndian.Uint16(decryptedLen)) + if payloadLen == 0 { + return nil, nil + } + + chunkEnd := payloadLen + AEADTagSize + if _, err := io.ReadFull(r.reader, r.buffer[:chunkEnd]); err != nil { + return nil, err + } + + decryptedPayload, err := r.cipher.Open(r.buffer[:0], r.nonce[:], r.buffer[:chunkEnd], nil) + if err != nil { + return nil, errors.New("failed to decrypt chunk payload").Base(err) + } + IncreaseNonce(r.nonce[:]) + + b := buf.New() + b.Write(decryptedPayload) + return buf.MultiBuffer{b}, nil +} + +type ClientRequestHeader struct { + Destination net.Destination + EarlyData []byte +} + +func ReadClientRequestHeader(conn io.Reader, reader *StreamReader) (*ClientRequestHeader, error) { + var fixedBuf [RequestHeaderFixedChunkLength + AEADTagSize]byte + if _, err := io.ReadFull(conn, fixedBuf[:]); err != nil { + return nil, err + } + + plainFixed, err := reader.cipher.Open(fixedBuf[:0], reader.Nonce(), fixedBuf[:], nil) + if err != nil { + return nil, errors.New("failed to decrypt client request header").Base(err) + } + IncreaseNonce(reader.Nonce()) + + if plainFixed[0] != HeaderTypeClient { + return nil, ErrBadHeaderType + } + + epoch := binary.BigEndian.Uint64(plainFixed[1:9]) + diff := int(math.Abs(float64(time.Now().Unix() - int64(epoch)))) + if diff > 30 { + return nil, ErrBadTimestamp + } + + varHeaderLen := int(binary.BigEndian.Uint16(plainFixed[9:11])) + if varHeaderLen == 0 { + return nil, ErrInvalidRequest + } + + var stackVarChunk [512]byte + var varChunkCipher []byte + needed := varHeaderLen + AEADTagSize + if needed <= len(stackVarChunk) { + varChunkCipher = stackVarChunk[:needed] + } else { + varChunkCipher = make([]byte, needed) + } + if _, err := io.ReadFull(conn, varChunkCipher); err != nil { + return nil, err + } + + plainVar, err := reader.cipher.Open(varChunkCipher[:0], reader.Nonce(), varChunkCipher, nil) + if err != nil { + return nil, errors.New("failed to decrypt variable request header").Base(err) + } + IncreaseNonce(reader.Nonce()) + + b := buf.New() + b.Write(plainVar) + defer b.Release() + + dest, err := ReadAddressPort(b) + if err != nil { + return nil, err + } + + var padLenBytes [2]byte + if _, err := b.Read(padLenBytes[:]); err != nil { + return nil, err + } + paddingLen := int(binary.BigEndian.Uint16(padLenBytes[:])) + if int(b.Len()) < paddingLen { + return nil, ErrNoPadding + } + if paddingLen > 0 { + b.Advance(int32(paddingLen)) + } + + var earlyData []byte + if b.Len() > 0 { + earlyData = make([]byte, b.Len()) + copy(earlyData, b.Bytes()) + } + + return &ClientRequestHeader{ + Destination: dest, + EarlyData: earlyData, + }, nil +} + +// ClientHandshake writes the full client request header to w +func ClientHandshake(w io.Writer, method *CipherMethod, pskList [][]byte, dest net.Destination, payload []byte) ([]byte, *StreamWriter, error) { + salt := make([]byte, method.KeySaltLength) + if _, err := io.ReadFull(rand.Reader, salt); err != nil { + return nil, nil, err + } + writer, err := WriteTCPRequest(w, method, pskList, dest, salt, payload) + if err != nil { + return nil, nil, err + } + return salt, writer.(*StreamWriter), nil +} + +// ClientVerifyServerResponse reads and verifies the server's handshake response +func ClientVerifyServerResponse(r io.Reader, method *CipherMethod, psk []byte, clientSalt []byte) (*StreamReader, []byte, error) { + reader, err := ReadTCPResponse(r, method, psk, clientSalt) + if err != nil { + return nil, nil, err + } + sr := reader.(*StreamReader) + var initialPayload []byte + if sr.cached > 0 { + initialPayload = make([]byte, sr.cached) + copy(initialPayload, sr.buffer[sr.offset:sr.offset+sr.cached]) + } + return sr, initialPayload, nil +} + +// WriteTCPRequest writes the Shadowsocks 2022 request header into w and returns a body writer. +func WriteTCPRequest(w io.Writer, method *CipherMethod, pskList [][]byte, dest net.Destination, clientSalt []byte, payload []byte) (buf.Writer, error) { + finalPSK := pskList[len(pskList)-1] + sessionKey := DeriveSessionSubKey(finalPSK, clientSalt, method.KeySaltLength) + aead, err := method.NewAEAD(sessionKey) + if err != nil { + return nil, err + } + + writer := NewStreamWriter(w, aead) + + handshakeBuf := buf.New() + defer handshakeBuf.Release() + + handshakeBuf.Write(clientSalt) + + if len(pskList) > 1 { + for i := 0; i < len(pskList)-1; i++ { + currPSK := pskList[i] + identitySubkey := DeriveIdentitySubKey(currPSK, clientSalt, method.KeySaltLength) + block, err := method.NewBlock(identitySubkey) + if err != nil { + return nil, err + } + nextPSK := pskList[i+1] + pskHash := DeriveUserPSKHash(nextPSK) + var encryptedEIH [AESBlockSize]byte + block.Encrypt(encryptedEIH[:], pskHash[:]) + handshakeBuf.Write(encryptedEIH[:]) + } + } + + payloadLen := len(payload) + var paddingLen int + if payloadLen < MaxPaddingLength { + paddingLen = mrand.Intn(MaxPaddingLength-payloadLen) + 1 + } + addrPortLen := AddrPortLength(dest) + varHeaderLen := addrPortLen + 2 + paddingLen + payloadLen + + var fixedHeaderPlaintext [RequestHeaderFixedChunkLength]byte + fixedHeaderPlaintext[0] = HeaderTypeClient + binary.BigEndian.PutUint64(fixedHeaderPlaintext[1:9], uint64(time.Now().Unix())) + binary.BigEndian.PutUint16(fixedHeaderPlaintext[9:11], uint16(varHeaderLen)) + + fixedChunk := writer.cipher.Seal(nil, writer.nonce[:], fixedHeaderPlaintext[:], nil) + IncreaseNonce(writer.nonce[:]) + handshakeBuf.Write(fixedChunk) + + varHeaderBuf := buf.New() + defer varHeaderBuf.Release() + + if err := WriteAddressPort(varHeaderBuf, dest); err != nil { + return nil, err + } + + var padLenBytes [2]byte + binary.BigEndian.PutUint16(padLenBytes[:], uint16(paddingLen)) + varHeaderBuf.Write(padLenBytes[:]) + + if paddingLen > 0 { + varHeaderBuf.Write(zeroPadding[:paddingLen]) + } + + if payloadLen > 0 { + varHeaderBuf.Write(payload) + } + + varChunk := writer.cipher.Seal(nil, writer.nonce[:], varHeaderBuf.Bytes(), nil) + IncreaseNonce(writer.nonce[:]) + handshakeBuf.Write(varChunk) + + if _, err := w.Write(handshakeBuf.Bytes()); err != nil { + return nil, err + } + + return writer, nil +} + +// ReadTCPResponse reads and verifies the server's handshake response and returns a reader for the stream. +func ReadTCPResponse(r io.Reader, method *CipherMethod, psk []byte, clientSalt []byte) (buf.Reader, error) { + var serverSalt [32]byte + serverSaltSlice := serverSalt[:method.KeySaltLength] + if _, err := io.ReadFull(r, serverSaltSlice); err != nil { + return nil, err + } + + sessionKey := DeriveSessionSubKey(psk, serverSaltSlice, method.KeySaltLength) + aead, err := method.NewAEAD(sessionKey) + if err != nil { + return nil, err + } + + reader := NewStreamReader(r, aead) + + fixedPlainLen := 1 + 8 + method.KeySaltLength + 2 + chunkCipherLen := fixedPlainLen + AEADTagSize + var chunkBuf [64]byte + chunkSlice := chunkBuf[:chunkCipherLen] + if _, err := io.ReadFull(r, chunkSlice); err != nil { + return nil, err + } + + decryptedFixed, err := reader.cipher.Open(chunkSlice[:0], reader.nonce[:], chunkSlice, nil) + if err != nil { + return nil, errors.New("failed to decrypt server response header").Base(err) + } + IncreaseNonce(reader.nonce[:]) + + if decryptedFixed[0] != HeaderTypeServer { + return nil, ErrBadHeaderType + } + + serverEpoch := binary.BigEndian.Uint64(decryptedFixed[1:9]) + diff := int(math.Abs(float64(time.Now().Unix() - int64(serverEpoch)))) + if diff > 30 { + return nil, ErrBadTimestamp + } + + echoedSalt := decryptedFixed[9 : 9+method.KeySaltLength] + for i := 0; i < method.KeySaltLength; i++ { + if echoedSalt[i] != clientSalt[i] { + return nil, errors.New("bad request salt") + } + } + + initialPayloadLen := int(binary.BigEndian.Uint16(decryptedFixed[9+method.KeySaltLength : 11+method.KeySaltLength])) + if initialPayloadLen > 0 { + initialCipherLen := initialPayloadLen + AEADTagSize + if _, err := io.ReadFull(r, reader.buffer[:initialCipherLen]); err != nil { + return nil, err + } + decryptedInitial, err := reader.cipher.Open(reader.buffer[:0], reader.nonce[:], reader.buffer[:initialCipherLen], nil) + if err != nil { + return nil, errors.New("failed to decrypt initial response payload").Base(err) + } + IncreaseNonce(reader.nonce[:]) + reader.cached = len(decryptedInitial) + reader.offset = 0 + } + + return reader, nil +} + +// WriteTCPResponse writes the server handshake response and returns a body writer for server stream. +func WriteTCPResponse(w io.Writer, method *CipherMethod, psk []byte, clientSalt []byte, initialPayload []byte) (buf.Writer, error) { + var serverSalt [32]byte + serverSaltSlice := serverSalt[:method.KeySaltLength] + if _, err := io.ReadFull(rand.Reader, serverSaltSlice); err != nil { + return nil, err + } + + respKey := DeriveSessionSubKey(psk, serverSaltSlice, method.KeySaltLength) + respAead, err := method.NewAEAD(respKey) + if err != nil { + return nil, err + } + writer := NewStreamWriter(w, respAead) + + respBuf := buf.New() + defer respBuf.Release() + + respBuf.Write(serverSaltSlice) + + var fixedRespPlain [1 + 8 + 32 + 2]byte + fixedRespSlice := fixedRespPlain[:1+8+method.KeySaltLength+2] + fixedRespSlice[0] = HeaderTypeServer + binary.BigEndian.PutUint64(fixedRespSlice[1:9], uint64(time.Now().Unix())) + copy(fixedRespSlice[9:9+method.KeySaltLength], clientSalt) + binary.BigEndian.PutUint16(fixedRespSlice[9+method.KeySaltLength:11+method.KeySaltLength], uint16(len(initialPayload))) + + fixedRespChunk := writer.cipher.Seal(nil, writer.nonce[:], fixedRespSlice, nil) + IncreaseNonce(writer.nonce[:]) + respBuf.Write(fixedRespChunk) + + if len(initialPayload) > 0 { + initialChunk := writer.cipher.Seal(nil, writer.nonce[:], initialPayload, nil) + IncreaseNonce(writer.nonce[:]) + respBuf.Write(initialChunk) + } + + if _, err := w.Write(respBuf.Bytes()); err != nil { + return nil, err + } + + return writer, nil +} diff --git a/testing/scenarios/shadowsocks_2022_test.go b/testing/scenarios/shadowsocks_2022_test.go index c5e0f8b3d..f456c5a3c 100644 --- a/testing/scenarios/shadowsocks_2022_test.go +++ b/testing/scenarios/shadowsocks_2022_test.go @@ -6,7 +6,6 @@ import ( "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" @@ -22,9 +21,19 @@ import ( "golang.org/x/sync/errgroup" ) +var ss2022Methods = []string{ + shadowsocks_2022.MethodAES128GCM, + shadowsocks_2022.MethodAES256GCM, + shadowsocks_2022.MethodChaCha20Poly1305, +} + func TestShadowsocks2022Tcp(t *testing.T) { - for _, method := range shadowaead_2022.List { - password := make([]byte, 32) + for _, method := range ss2022Methods { + keySize := 32 + if method == shadowsocks_2022.MethodAES128GCM { + keySize = 16 + } + password := make([]byte, keySize) rand.Read(password) t.Run(method, func(t *testing.T) { testShadowsocks2022Tcp(t, method, base64.StdEncoding.EncodeToString(password)) @@ -33,21 +42,21 @@ func TestShadowsocks2022Tcp(t *testing.T) { } func TestShadowsocks2022UdpAES128(t *testing.T) { - password := make([]byte, 32) + password := make([]byte, 16) rand.Read(password) - testShadowsocks2022Udp(t, shadowaead_2022.List[0], base64.StdEncoding.EncodeToString(password)) + testShadowsocks2022Udp(t, shadowsocks_2022.MethodAES128GCM, 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)) + testShadowsocks2022Udp(t, shadowsocks_2022.MethodAES256GCM, 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)) + testShadowsocks2022Udp(t, shadowsocks_2022.MethodChaCha20Poly1305, base64.StdEncoding.EncodeToString(password)) } func testShadowsocks2022Tcp(t *testing.T, method string, password string) {