This commit is contained in:
Fangliding
2026-09-26 16:38:19 +08:00
parent 643fa66adc
commit fe1c2f9bb8
12 changed files with 125 additions and 572 deletions
+34 -54
View File
@@ -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
}
+10 -14
View File
@@ -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)
+3 -3
View File
@@ -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
+4 -4
View File
@@ -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
}
+2 -2
View File
@@ -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{
+2 -4
View File
@@ -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 {
+3 -7
View File
@@ -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
+42
View File
@@ -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
}
+4 -4
View File
@@ -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) {
@@ -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
}
+11 -443
View File
@@ -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")
}
}
+10 -21
View File
@@ -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)