mirror of
https://github.com/XTLS/Xray-core.git
synced 2026-09-24 18:10:32 +00:00
Compare commits
| Author | SHA1 | Date | |
|---|---|---|---|
|
|
3b905381fb | ||
|
|
ea3275a77b |
@@ -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
|
||||
}
|
||||
@@ -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
|
||||
}
|
||||
@@ -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
|
||||
}
|
||||
@@ -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())
|
||||
}
|
||||
@@ -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) {
|
||||
}
|
||||
@@ -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
|
||||
}
|
||||
@@ -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
|
||||
}
|
||||
@@ -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
|
||||
}
|
||||
@@ -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
|
||||
|
||||
@@ -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=
|
||||
|
||||
@@ -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 {
|
||||
|
||||
@@ -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")
|
||||
}
|
||||
@@ -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),
|
||||
}
|
||||
}
|
||||
|
||||
+208
-126
@@ -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 {
|
||||
@@ -70,114 +85,181 @@ func (i *Inbound) Process(ctx context.Context, network net.Network, connection s
|
||||
inbound := session.InboundFromContext(ctx)
|
||||
inbound.Name = "shadowsocks-2022"
|
||||
inbound.CanSpliceCopy = 3
|
||||
|
||||
var metadata M.Metadata
|
||||
if inbound.Source.IsValid() {
|
||||
metadata.Source = M.ParseSocksaddr(inbound.Source.NetAddr())
|
||||
}
|
||||
|
||||
ctx = session.ContextWithDispatcher(ctx, dispatcher)
|
||||
inbound.User = i.user
|
||||
|
||||
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()
|
||||
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()
|
||||
}
|
||||
|
||||
var salt [32]byte
|
||||
saltSlice := salt[:i.method.KeySaltLength]
|
||||
if _, err := io.ReadFull(conn, saltSlice); err != nil {
|
||||
return err
|
||||
}
|
||||
|
||||
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
|
||||
}
|
||||
|
||||
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
|
||||
})
|
||||
}()
|
||||
|
||||
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
|
||||
}
|
||||
|
||||
@@ -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,319 @@ 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)
|
||||
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)
|
||||
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)
|
||||
}
|
||||
|
||||
@@ -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,206 @@ 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)
|
||||
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)
|
||||
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())
|
||||
}
|
||||
|
||||
@@ -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
|
||||
}
|
||||
@@ -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)
|
||||
}
|
||||
|
||||
@@ -0,0 +1,516 @@
|
||||
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
|
||||
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
|
||||
|
||||
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 (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 (c *UDPCodec) EncodeServerPacket(clientSessionID uint64, dest net.Destination, payload []byte) ([]byte, error) {
|
||||
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)
|
||||
}
|
||||
|
||||
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
|
||||
}
|
||||
}
|
||||
@@ -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() {}
|
||||
@@ -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)
|
||||
}
|
||||
@@ -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
|
||||
}
|
||||
|
||||
@@ -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")
|
||||
}
|
||||
}
|
||||
@@ -0,0 +1,540 @@
|
||||
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")
|
||||
}
|
||||
|
||||
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
|
||||
}
|
||||
@@ -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) {
|
||||
|
||||
Reference in New Issue
Block a user