From fe1c2f9bb895f1c7f1ab72f894c83d24b8ebf02f Mon Sep 17 00:00:00 2001 From: Fangliding Date: Sat, 26 Sep 2026 16:34:04 +0800 Subject: [PATCH] refine --- infra/conf/shadowsocks.go | 88 ++-- proxy/shadowsocks_2022/cipher.go | 24 +- proxy/shadowsocks_2022/inbound.go | 6 +- proxy/shadowsocks_2022/inbound_multi.go | 8 +- proxy/shadowsocks_2022/inbound_relay.go | 4 +- proxy/shadowsocks_2022/outbound.go | 6 +- proxy/shadowsocks_2022/packet.go | 10 +- proxy/shadowsocks_2022/relay_test.go | 42 ++ proxy/shadowsocks_2022/replay.go | 8 +- proxy/shadowsocks_2022/shadowsocks_2022.go | 16 - .../shadowsocks_2022/shadowsocks_2022_test.go | 454 +----------------- proxy/shadowsocks_2022/stream.go | 31 +- 12 files changed, 125 insertions(+), 572 deletions(-) diff --git a/infra/conf/shadowsocks.go b/infra/conf/shadowsocks.go index c4e7790e2..9032dad3f 100644 --- a/infra/conf/shadowsocks.go +++ b/infra/conf/shadowsocks.go @@ -53,7 +53,7 @@ func (v *ShadowsocksServerConfig) Build() (proto.Message, error) { v.Users = v.Clients } - if shadowsocks_2022.IsSupportedMethod(v.Cipher) { + if _, err := shadowsocks_2022.GetCipherMethod(v.Cipher); err == nil { return buildShadowsocks2022(v) } @@ -212,63 +212,43 @@ func (v *ShadowsocksClientConfig) Build() (proto.Message, error) { return nil, errors.New(`Shadowsocks settings: "servers" should have one and only one member. Multiple endpoints in "servers" should use multiple Shadowsocks outbounds and routing balancer instead`) } - if len(v.Servers) == 1 { - server := v.Servers[0] - if shadowsocks_2022.IsSupportedMethod(server.Cipher) { - if server.Address == nil { - return nil, errors.New("Shadowsocks server address is not set.") - } - if server.Port == 0 { - return nil, errors.New("Invalid Shadowsocks port.") - } - if server.Password == "" { - return nil, errors.New("Shadowsocks password is not specified.") - } - - config := new(shadowsocks_2022.ClientConfig) - config.Address = server.Address.Build() - config.Port = uint32(server.Port) - config.Method = server.Cipher - config.Key = server.Password - return config, nil - } + server := v.Servers[0] + if server.Address == nil { + return nil, errors.New("Shadowsocks server address is not set.") + } + if server.Port == 0 { + return nil, errors.New("Invalid Shadowsocks port.") + } + if server.Password == "" { + return nil, errors.New("Shadowsocks password is not specified.") } + if _, err := shadowsocks_2022.GetCipherMethod(server.Cipher); err == nil { + config := new(shadowsocks_2022.ClientConfig) + config.Address = server.Address.Build() + config.Port = uint32(server.Port) + config.Method = server.Cipher + config.Key = server.Password + return config, nil + } config := new(shadowsocks.ClientConfig) - for _, server := range v.Servers { - if shadowsocks_2022.IsSupportedMethod(server.Cipher) { - return nil, errors.New("Shadowsocks 2022 accept no multi servers") - } - if server.Address == nil { - return nil, errors.New("Shadowsocks server address is not set.") - } - if server.Port == 0 { - return nil, errors.New("Invalid Shadowsocks port.") - } - if server.Password == "" { - return nil, errors.New("Shadowsocks password is not specified.") - } - account := &shadowsocks.Account{ - Password: server.Password, - } - account.CipherType = cipherFromString(server.Cipher) - if account.CipherType == shadowsocks.CipherType_UNKNOWN { - return nil, errors.New("unknown cipher method: ", server.Cipher) - } - - ss := &protocol.ServerEndpoint{ - Address: server.Address.Build(), - Port: uint32(server.Port), - User: &protocol.User{ - Level: uint32(server.Level), - Email: server.Email, - Account: serial.ToTypedMessage(account), - }, - } - - config.Server = ss - break + account := &shadowsocks.Account{ + Password: server.Password, } + account.CipherType = cipherFromString(server.Cipher) + if account.CipherType == shadowsocks.CipherType_UNKNOWN { + return nil, errors.New("unknown cipher method: ", server.Cipher) + } + ss := &protocol.ServerEndpoint{ + Address: server.Address.Build(), + Port: uint32(server.Port), + User: &protocol.User{ + Level: uint32(server.Level), + Email: server.Email, + Account: serial.ToTypedMessage(account), + }, + } + config.Server = ss return config, nil } diff --git a/proxy/shadowsocks_2022/cipher.go b/proxy/shadowsocks_2022/cipher.go index dde381c92..311b9bca9 100644 --- a/proxy/shadowsocks_2022/cipher.go +++ b/proxy/shadowsocks_2022/cipher.go @@ -4,6 +4,7 @@ import ( "crypto/aes" "crypto/cipher" "errors" + "strings" "golang.org/x/crypto/chacha20poly1305" ) @@ -14,23 +15,18 @@ type CipherMethod struct { 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} -) +var methods = map[string]*CipherMethod{ + MethodAES128GCM: {Name: MethodAES128GCM, KeySaltLength: 16, IsChaCha: false}, + MethodAES256GCM: {Name: MethodAES256GCM, KeySaltLength: 32, IsChaCha: false}, + MethodChaCha20Poly1305: {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") + name = strings.ToLower(name) + if m, ok := methods[name]; ok { + return m, nil } + return nil, errors.New("unknown shadowsocks 2022 method") } // NewAEAD creates standard stream AEAD cipher instance (AES-GCM or ChaCha20-Poly1305) diff --git a/proxy/shadowsocks_2022/inbound.go b/proxy/shadowsocks_2022/inbound.go index 19c768ee7..eaa637d30 100644 --- a/proxy/shadowsocks_2022/inbound.go +++ b/proxy/shadowsocks_2022/inbound.go @@ -98,7 +98,7 @@ func (i *Inbound) processTCP(ctx context.Context, conn net.Conn, dispatcher rout 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() + return errors.New("unable to set read deadline").Base(err) } var salt [32]byte @@ -123,7 +123,7 @@ func (i *Inbound) processTCP(ctx context.Context, conn net.Conn, dispatcher rout if err != nil { return err } - _ = conn.SetReadDeadline(time.Time{}) + conn.SetReadDeadline(time.Time{}) dest := reqHeader.Destination writer, err := WriteTCPResponse(conn, i.method, i.psk, saltSlice, nil) @@ -243,7 +243,7 @@ func (i *Inbound) processUDP(ctx context.Context, conn stat.Connection, dispatch } cEntry.timer.Update() for _, rb := range resMb { - encPacket, err := i.udpCodec.EncodePacket(sessID, dest, rb.Bytes()) + encPacket, err := i.udpCodec.EncodeServerPacket(sessID, dest, rb.Bytes()) rb.Release() if err != nil { continue diff --git a/proxy/shadowsocks_2022/inbound_multi.go b/proxy/shadowsocks_2022/inbound_multi.go index ba54e4a21..6a8251bb8 100644 --- a/proxy/shadowsocks_2022/inbound_multi.go +++ b/proxy/shadowsocks_2022/inbound_multi.go @@ -204,7 +204,7 @@ func (i *MultiUserInbound) processTCP(ctx context.Context, conn net.Conn, dispat 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() + return errors.New("unable to set read deadline").Base(err) } // 1. Read Request Salt (16 or 32 bytes) @@ -255,7 +255,7 @@ func (i *MultiUserInbound) processTCP(ctx context.Context, conn net.Conn, dispat if err != nil { return err } - _ = conn.SetReadDeadline(time.Time{}) + conn.SetReadDeadline(time.Time{}) dest := reqHeader.Destination // 6. Send Server Response Handshake @@ -343,7 +343,7 @@ func (i *MultiUserInbound) processUDP(ctx context.Context, conn stat.Connection, packetID := binary.BigEndian.Uint64(rawHeader[8:16]) // Replay protection & session lookup - sessionItem, _ := i.udpSessions.GetOrCreate(sessionID) + sessionItem := i.udpSessions.GetOrCreate(sessionID) sessionItem.Lock() if !sessionItem.Window.Check(packetID) { @@ -503,7 +503,7 @@ func (i *MultiUserInbound) processUDP(ctx context.Context, conn stat.Connection, } func (i *MultiUserInbound) encodeServerUDPPacket(clientSessionID uint64, userPSK []byte, dest net.Destination, payload []byte) ([]byte, error) { - sessionItem, _ := i.udpSessions.GetOrCreate(clientSessionID) + sessionItem := i.udpSessions.GetOrCreate(clientSessionID) if err := sessionItem.EnsureServerState(i.method, i.udpMasterCipher, nil, userPSK); err != nil { return nil, err } diff --git a/proxy/shadowsocks_2022/inbound_relay.go b/proxy/shadowsocks_2022/inbound_relay.go index 53eea1b4d..8f29db12e 100644 --- a/proxy/shadowsocks_2022/inbound_relay.go +++ b/proxy/shadowsocks_2022/inbound_relay.go @@ -136,7 +136,7 @@ func (i *RelayInbound) processTCP(ctx context.Context, conn net.Conn, dispatcher 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() + return errors.New("unable to set read deadline").Base(err) } // Read Salt + Outer EIH @@ -163,7 +163,7 @@ func (i *RelayInbound) processTCP(ctx context.Context, conn net.Conn, dispatcher if !ok { return ErrInvalidRequest } - _ = conn.SetReadDeadline(time.Time{}) + conn.SetReadDeadline(time.Time{}) inbound := session.InboundFromContext(ctx) inbound.User = &protocol.MemoryUser{ diff --git a/proxy/shadowsocks_2022/outbound.go b/proxy/shadowsocks_2022/outbound.go index f42590d81..71d21475d 100644 --- a/proxy/shadowsocks_2022/outbound.go +++ b/proxy/shadowsocks_2022/outbound.go @@ -27,7 +27,6 @@ func init() { } type Outbound struct { - ctx context.Context server net.Destination method *CipherMethod pskList [][]byte @@ -55,7 +54,6 @@ func NewClient(ctx context.Context, config *ClientConfig) (*Outbound, error) { v := core.MustFromContext(ctx) return &Outbound{ - ctx: ctx, server: net.Destination{ Address: config.Address.AsAddress(), Port: net.Port(config.Port), @@ -94,7 +92,7 @@ func (o *Outbound) Process(ctx context.Context, link *transport.Link, dialer int conn = rawConn return nil }); err != nil { - return errors.New("failed to find an available destination").Base(err).AtWarning() + return errors.New("failed to find an available destination").Base(err) } defer conn.Close() @@ -135,7 +133,7 @@ func (o *Outbound) Process(ctx context.Context, link *transport.Link, dialer int } 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() + return errors.New("failed to write A request payload").Base(err) } if err := bufferedWriter.SetBuffered(false); err != nil { diff --git a/proxy/shadowsocks_2022/packet.go b/proxy/shadowsocks_2022/packet.go index 7a0429fdf..f0d1daba7 100644 --- a/proxy/shadowsocks_2022/packet.go +++ b/proxy/shadowsocks_2022/packet.go @@ -281,7 +281,7 @@ func (c *UDPCodec) DecodePacket(data []byte) (DecodedUDPPacket, error) { packetID := binary.BigEndian.Uint64(plain[8:16]) if c.sessions != nil { - sessionItem, _ := c.sessions.GetOrCreate(sessionID) + sessionItem := c.sessions.GetOrCreate(sessionID) sessionItem.Lock() if !sessionItem.Window.CheckAndAdd(packetID) { sessionItem.Unlock() @@ -303,7 +303,7 @@ func (c *UDPCodec) DecodePacket(data []byte) (DecodedUDPPacket, error) { var sessionItem *ServerUDPSession if c.sessions != nil { - sessionItem, _ = c.sessions.GetOrCreate(sessionID) + sessionItem = c.sessions.GetOrCreate(sessionID) sessionItem.Lock() if !sessionItem.Window.Check(packetID) { sessionItem.Unlock() @@ -444,17 +444,13 @@ func (s *ServerUDPSession) EncodeServerPacket(method *CipherMethod, clientSessio } func (c *UDPCodec) EncodeServerPacket(clientSessionID uint64, dest net.Destination, payload []byte) ([]byte, error) { - sessionItem, _ := c.sessions.GetOrCreate(clientSessionID) + 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 diff --git a/proxy/shadowsocks_2022/relay_test.go b/proxy/shadowsocks_2022/relay_test.go index bb8a3277e..26834ece5 100644 --- a/proxy/shadowsocks_2022/relay_test.go +++ b/proxy/shadowsocks_2022/relay_test.go @@ -10,9 +10,12 @@ import ( "testing" "time" + "errors" + "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/features/routing" . "github.com/xtls/xray-core/proxy/shadowsocks_2022" "github.com/xtls/xray-core/transport" "lukechampine.com/blake3" @@ -228,3 +231,42 @@ func (w *customWriter) Close() error { } func (w *customWriter) Interrupt() {} + +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 +} diff --git a/proxy/shadowsocks_2022/replay.go b/proxy/shadowsocks_2022/replay.go index 86a47d8c4..eb59bae0b 100644 --- a/proxy/shadowsocks_2022/replay.go +++ b/proxy/shadowsocks_2022/replay.go @@ -114,11 +114,11 @@ func NewUDPSessionManager(timeout time.Duration) *UDPSessionManager { } } -func (m *UDPSessionManager) GetOrCreate(sessionID uint64) (*ServerUDPSession, bool) { +func (m *UDPSessionManager) GetOrCreate(sessionID uint64) *ServerUDPSession { now := time.Now().Unix() if s, ok := m.sessions.Load(sessionID); ok { s.LastActive.Store(now) - return s, true + return s } s := &ServerUDPSession{ @@ -129,7 +129,7 @@ func (m *UDPSessionManager) GetOrCreate(sessionID uint64) (*ServerUDPSession, bo actual, loaded := m.sessions.LoadOrStore(sessionID, s) if loaded { actual.LastActive.Store(now) - return actual, true + return actual } // Trigger cleanup if at least 30 seconds have passed since last cleanup @@ -138,7 +138,7 @@ func (m *UDPSessionManager) GetOrCreate(sessionID uint64) (*ServerUDPSession, bo go m.cleanup(now) } - return s, false + return s } func (m *UDPSessionManager) cleanup(now int64) { diff --git a/proxy/shadowsocks_2022/shadowsocks_2022.go b/proxy/shadowsocks_2022/shadowsocks_2022.go index 697cf5955..505f33c2d 100644 --- a/proxy/shadowsocks_2022/shadowsocks_2022.go +++ b/proxy/shadowsocks_2022/shadowsocks_2022.go @@ -2,7 +2,6 @@ package shadowsocks_2022 import ( "context" - "strings" "sync" "github.com/xtls/xray-core/common/errors" @@ -38,12 +37,6 @@ const ( MethodChaCha20Poly1305 = "2022-blake3-chacha20-poly1305" ) -var List = []string{ - MethodAES128GCM, - MethodAES256GCM, - MethodChaCha20Poly1305, -} - var ( ErrBadKey = errors.New("bad key") ErrBadHeaderType = errors.New("bad header type") @@ -55,12 +48,3 @@ var ( ErrNoPadding = errors.New("bad request: missing payload or padding") ErrInvalidRequest = errors.New("invalid request") ) - -func IsSupportedMethod(method string) bool { - for _, m := range List { - if strings.EqualFold(m, method) { - return true - } - } - return false -} diff --git a/proxy/shadowsocks_2022/shadowsocks_2022_test.go b/proxy/shadowsocks_2022/shadowsocks_2022_test.go index 0f8df2c49..f6f930a27 100644 --- a/proxy/shadowsocks_2022/shadowsocks_2022_test.go +++ b/proxy/shadowsocks_2022/shadowsocks_2022_test.go @@ -14,18 +14,13 @@ import ( "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 { @@ -87,24 +82,7 @@ func TestKDF(t *testing.T) { } } -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 +func TestSlidingWindow(t *testing.T) { var window SlidingWindow if !window.Check(1) { t.Fatal("packet 1 should be accepted") @@ -139,7 +117,7 @@ func TestReplayFilter(t *testing.T) { } } -func TestTCPStreamAndHandshake(t *testing.T) { +func TestTCPStream(t *testing.T) { methods := []struct { name string keySize int @@ -303,6 +281,9 @@ func TestUDPCodec(t *testing.T) { common.Must(err) defer pktBuf.Release() + rawCopy := make([]byte, pktBuf.Len()) + copy(rawCopy, pktBuf.Bytes()) + decoded, err := serverCodec.DecodePacket(pktBuf.Bytes()) common.Must(err) @@ -315,6 +296,12 @@ func TestUDPCodec(t *testing.T) { if !bytes.Equal(decoded.Payload, payload) { t.Errorf("payload mismatch: got %s, want %s", decoded.Payload, payload) } + + // 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) + } }) } } @@ -373,422 +360,3 @@ func TestMultiUserManager(t *testing.T) { t.Fatal("user1 should have been removed") } } - -type dummyDispatcher struct { - onDispatch func(ctx context.Context, dest net.Destination) (*transport.Link, error) -} - -func (d *dummyDispatcher) Dispatch(ctx context.Context, dest net.Destination) (*transport.Link, error) { - if d.onDispatch != nil { - return d.onDispatch(ctx, dest) - } - return nil, errors.New("not handled") -} - -func (d *dummyDispatcher) DispatchLink(ctx context.Context, dest net.Destination, link *transport.Link) error { - return nil -} - -func (d *dummyDispatcher) Start() error { return nil } -func (d *dummyDispatcher) Close() error { return nil } -func (d *dummyDispatcher) Type() interface{} { return routing.DispatcherType() } - -type dummyStatConn struct { - gonet.Conn -} - -func (c *dummyStatConn) ReadMultiBuffer() (buf.MultiBuffer, error) { - b := buf.New() - _, err := b.ReadFrom(c.Conn) - return buf.MultiBuffer{b}, err -} - -func (c *dummyStatConn) WriteMultiBuffer(mb buf.MultiBuffer) error { - defer buf.ReleaseMulti(mb) - for _, b := range mb { - if _, err := c.Conn.Write(b.Bytes()); err != nil { - return err - } - } - return nil -} - -func TestMultiUserTCPConnection(t *testing.T) { - masterKey := generateRandomKey(16) - userKey1 := generateRandomKey(16) - userKey2 := generateRandomKey(16) - - config := &MultiUserServerConfig{ - Method: MethodAES128GCM, - Key: masterKey, - Users: []*protocol.User{ - { - Email: "user1@example.com", - Account: serial.ToTypedMessage(&Account{Key: userKey1}), - }, - { - Email: "user2@example.com", - Account: serial.ToTypedMessage(&Account{Key: userKey2}), - }, - }, - } - - testCtx := newTestContext() - inbound, err := NewMultiServer(testCtx, config) - common.Must(err) - - clientConn, serverConn := gonet.Pipe() - defer clientConn.Close() - defer serverConn.Close() - - dest := net.TCPDestination(net.LocalHostIP, 443) - method, err := GetCipherMethod(MethodAES128GCM) - common.Must(err) - - masterRaw, _ := base64.StdEncoding.DecodeString(masterKey) - user2Raw, _ := base64.StdEncoding.DecodeString(userKey2) - clientPSKList := [][]byte{masterRaw, user2Raw} - - dispatchedUserChan := make(chan string, 1) - - disp := &dummyDispatcher{ - onDispatch: func(ctx context.Context, d net.Destination) (*transport.Link, error) { - inbound := session.InboundFromContext(ctx) - if inbound != nil && inbound.User != nil { - dispatchedUserChan <- inbound.User.Email - } - link := &transport.Link{ - Reader: buf.NewReader(bytes.NewReader(nil)), - Writer: buf.Discard, - } - return link, nil - }, - } - - go func() { - _ = inbound.Process(testCtx, net.Network_TCP, &dummyStatConn{Conn: serverConn}, disp) - }() - - clientSalt, writer, err := ClientHandshake(clientConn, method, clientPSKList, dest, []byte("ping")) - common.Must(err) - - reader, _, err := ClientVerifyServerResponse(clientConn, method, user2Raw, clientSalt) - common.Must(err) - _ = writer - _ = reader - - select { - case email := <-dispatchedUserChan: - if email != "user2@example.com" { - t.Fatalf("expected user2@example.com, got %s", email) - } - case <-time.After(2 * time.Second): - t.Fatal("timeout waiting for dispatched user") - } -} - -func TestUDPReaderWriter(t *testing.T) { - for _, methodName := range []string{MethodAES128GCM, MethodAES256GCM, MethodChaCha20Poly1305} { - t.Run(methodName, func(t *testing.T) { - method, err := GetCipherMethod(methodName) - common.Must(err) - rawPSK := make([]byte, method.KeySaltLength) - _, _ = rand.Read(rawPSK) - - clientCodec, err := NewUDPPacketCodec(method, rawPSK) - common.Must(err) - serverCodec, err := NewUDPServerCodec(method, rawPSK, time.Minute) - common.Must(err) - - dest := net.UDPDestination(net.LocalHostIP, 53) - - // Client to Server - clientPacketBuf, err := clientCodec.EncodeClientPacket(dest, []byte("hello dns")) - common.Must(err) - defer clientPacketBuf.Release() - - serverDecoded, err := serverCodec.DecodePacket(clientPacketBuf.Bytes()) - common.Must(err) - if string(serverDecoded.Payload) != "hello dns" { - t.Fatalf("unexpected server decoded payload: %s", string(serverDecoded.Payload)) - } - - // Server to Client - serverPacket, err := serverCodec.EncodePacket(serverDecoded.SessionID, dest, []byte("dns response")) - common.Must(err) - - clientDecoded, err := clientCodec.DecodePacket(serverPacket) - common.Must(err) - if string(clientDecoded.Payload) != "dns response" { - t.Fatalf("unexpected client decoded payload: %s", string(clientDecoded.Payload)) - } - - // Test UDPWriter and UDPReader pipeline - pipeR, pipeW := gonet.Pipe() - defer pipeR.Close() - defer pipeW.Close() - - writer := &UDPWriter{ - Writer: pipeW, - Destination: dest, - Codec: clientCodec, - } - reader := &UDPReader{ - Reader: pipeR, - Codec: clientCodec, - } - - go func() { - // Simulate server echoing back as server response - buf := make([]byte, 2048) - n, err := pipeR.Read(buf) - if err != nil { - return - } - dec, err := serverCodec.DecodePacket(buf[:n]) - if err != nil { - return - } - resp, err := serverCodec.EncodePacket(dec.SessionID, dest, dec.Payload) - if err != nil { - return - } - _, _ = pipeW.Write(resp) - }() - - b := buf.New() - b.WriteString("piped udp packet") - common.Must(writer.WriteMultiBuffer(buf.MultiBuffer{b})) - - received, err := reader.ReadMultiBuffer() - common.Must(err) - if received[0].String() != "piped udp packet" { - t.Fatalf("expected 'piped udp packet', got '%s'", received[0].String()) - } - }) - } -} - -func TestTCPRequestResponse(t *testing.T) { - for _, methodName := range []string{MethodAES128GCM, MethodAES256GCM, MethodChaCha20Poly1305} { - t.Run(methodName, func(t *testing.T) { - method, err := GetCipherMethod(methodName) - common.Must(err) - rawPSK := make([]byte, method.KeySaltLength) - _, _ = rand.Read(rawPSK) - - clientConn, serverConn := gonet.Pipe() - defer clientConn.Close() - defer serverConn.Close() - - clientSalt := make([]byte, method.KeySaltLength) - _, _ = rand.Read(clientSalt) - dest := net.TCPDestination(net.LocalHostIP, 80) - - go func() { - // Server side: read handshake and verify clientSalt - salt := make([]byte, method.KeySaltLength) - if _, err := io.ReadFull(serverConn, salt); err != nil { - t.Errorf("server read salt error: %v", err) - return - } - sessionKey := DeriveSessionSubKey(rawPSK, salt, method.KeySaltLength) - aead, err := method.NewAEAD(sessionKey) - if err != nil { - t.Errorf("server AEAD error: %v", err) - return - } - sReader := NewStreamReader(serverConn, aead) - var fixedBuf [RequestHeaderFixedChunkLength + AEADTagSize]byte - if _, err := io.ReadFull(serverConn, fixedBuf[:]); err != nil { - t.Errorf("server read fixed error: %v", err) - return - } - plainFixed, err := aead.Open(fixedBuf[:0], sReader.Nonce(), fixedBuf[:], nil) - if err != nil { - t.Errorf("server decrypt fixed error: %v", err) - return - } - IncreaseNonce(sReader.Nonce()) - - varLen := int(binary.BigEndian.Uint16(plainFixed[9:11])) - varBuf := make([]byte, varLen+AEADTagSize) - if _, err := io.ReadFull(serverConn, varBuf); err != nil { - t.Errorf("server read var error: %v", err) - return - } - plainVar, err := aead.Open(varBuf[:0], sReader.Nonce(), varBuf, nil) - if err != nil { - t.Errorf("server decrypt var error: %v", err) - return - } - IncreaseNonce(sReader.Nonce()) - - vBuf := buf.New() - vBuf.Write(plainVar) - receivedDest, err := ReadAddressPort(vBuf) - if err != nil || receivedDest != dest { - t.Errorf("dest mismatch: %v vs %v, err: %v", receivedDest, dest, err) - return - } - - // Echo client salt back to client using WriteTCPResponse - sWriter, err := WriteTCPResponse(serverConn, method, rawPSK, salt, []byte("early-reply")) - if err != nil { - t.Errorf("server response error: %v", err) - return - } - _ = sWriter - }() - - bodyWriter, err := WriteTCPRequest(clientConn, method, [][]byte{rawPSK}, dest, clientSalt, nil) - common.Must(err) - _ = bodyWriter - - responseReader, err := ReadTCPResponse(clientConn, method, rawPSK, clientSalt) - common.Must(err) - - mb, err := responseReader.ReadMultiBuffer() - common.Must(err) - if mb[0].String() != "early-reply" { - t.Fatalf("expected early-reply, got %s", mb[0].String()) - } - }) - } -} - -func TestUDPReplayProtection(t *testing.T) { - for _, methodName := range []string{MethodAES128GCM, MethodAES256GCM, MethodChaCha20Poly1305} { - t.Run(methodName, func(t *testing.T) { - method, err := GetCipherMethod(methodName) - common.Must(err) - rawPSK := make([]byte, method.KeySaltLength) - _, _ = rand.Read(rawPSK) - - clientCodec, err := NewUDPPacketCodec(method, rawPSK) - common.Must(err) - serverCodec, err := NewUDPServerCodec(method, rawPSK, time.Minute) - common.Must(err) - - dest := net.UDPDestination(net.LocalHostIP, 53) - pktBuf, err := clientCodec.EncodeClientPacket(dest, []byte("dns 1")) - common.Must(err) - defer pktBuf.Release() - - rawCopy := make([]byte, pktBuf.Len()) - copy(rawCopy, pktBuf.Bytes()) - - // First decode should succeed - _, err = serverCodec.DecodePacket(pktBuf.Bytes()) - if err != nil { - t.Fatalf("first decode failed: %v", err) - } - - // Replay same packet wire bytes should fail with ErrPacketIdNotUnique - _, err = serverCodec.DecodePacket(rawCopy) - if err != ErrPacketIdNotUnique { - t.Fatalf("expected ErrPacketIdNotUnique on replay, got: %v", err) - } - }) - } -} - -func TestServerUDPSessionStabilityAndMonotonicPacketID(t *testing.T) { - for _, methodName := range []string{MethodAES128GCM, MethodAES256GCM, MethodChaCha20Poly1305} { - t.Run(methodName, func(t *testing.T) { - method, err := GetCipherMethod(methodName) - common.Must(err) - rawPSK := make([]byte, method.KeySaltLength) - _, _ = rand.Read(rawPSK) - - clientCodec, err := NewUDPPacketCodec(method, rawPSK) - common.Must(err) - serverCodec, err := NewUDPServerCodec(method, rawPSK, time.Minute) - common.Must(err) - - dest := net.UDPDestination(net.LocalHostIP, 53) - - // Client sends packet 1 - pkt1, err := clientCodec.EncodeClientPacket(dest, []byte("request 1")) - common.Must(err) - defer pkt1.Release() - - dec1, err := serverCodec.DecodePacket(pkt1.Bytes()) - common.Must(err) - - // Server sends response 1 - resp1, err := serverCodec.EncodePacket(dec1.SessionID, dest, []byte("response 1")) - common.Must(err) - - // Server sends response 2 to the same client session - resp2, err := serverCodec.EncodePacket(dec1.SessionID, dest, []byte("response 2")) - common.Must(err) - - // Decode both on client - cDec1, err := clientCodec.DecodePacket(resp1) - common.Must(err) - cDec2, err := clientCodec.DecodePacket(resp2) - common.Must(err) - - if cDec1.SessionID != cDec2.SessionID { - t.Fatalf("expected stable server session ID, got %d and %d", cDec1.SessionID, cDec2.SessionID) - } - if cDec2.PacketID <= cDec1.PacketID { - t.Fatalf("expected monotonically increasing packet ID, got %d then %d", cDec1.PacketID, cDec2.PacketID) - } - if string(cDec1.Payload) != "response 1" || string(cDec2.Payload) != "response 2" { - t.Fatalf("payload mismatch") - } - }) - } -} - -type mockDialer struct { - dial func(ctx context.Context, dest net.Destination) (stat.Connection, error) -} - -func (d *mockDialer) Dial(ctx context.Context, dest net.Destination) (stat.Connection, error) { - if d.dial != nil { - return d.dial(ctx, dest) - } - c1, c2 := gonet.Pipe() - _ = c2.Close() - return &dummyStatConn{Conn: c1}, nil -} - -func (d *mockDialer) DestIpAddress() net.IP { - return net.IP{127, 0, 0, 1} -} - -func (d *mockDialer) SetOutboundGateway(ctx context.Context, ob *session.Outbound) {} - -func TestOutboundProcess(t *testing.T) { - testCtx := newTestContext() - key := generateRandomKey(16) - clientConfig := &ClientConfig{ - Address: &net.IPOrDomain{Address: &net.IPOrDomain_Ip{Ip: []byte{127, 0, 0, 1}}}, - Port: 1080, - Method: MethodAES128GCM, - Key: key, - } - - outbound, err := NewClient(testCtx, clientConfig) - common.Must(err) - - ctx := session.ContextWithOutbounds(testCtx, []*session.Outbound{ - { - Target: net.TCPDestination(net.LocalHostIP, 80), - }, - }) - - link := &transport.Link{ - Reader: buf.NewReader(bytes.NewReader(nil)), - Writer: buf.Discard, - } - - dialer := &mockDialer{} - err = outbound.Process(ctx, link, dialer) - if err == nil { - t.Fatal("expected error from closed mock dialer pipe, got nil") - } -} diff --git a/proxy/shadowsocks_2022/stream.go b/proxy/shadowsocks_2022/stream.go index 5d4a3c99f..b515df2da 100644 --- a/proxy/shadowsocks_2022/stream.go +++ b/proxy/shadowsocks_2022/stream.go @@ -81,10 +81,6 @@ 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 { @@ -152,10 +148,6 @@ 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]) @@ -367,20 +359,17 @@ func WriteTCPRequest(w io.Writer, method *CipherMethod, pskList [][]byte, dest n 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[:]) + for i, currPSK := range pskList[:len(pskList)-1] { + 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)