mirror of
https://github.com/XTLS/Xray-core.git
synced 2026-09-27 09:28:01 +03:00
refine
This commit is contained in:
+34
-54
@@ -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
|
||||
}
|
||||
|
||||
@@ -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)
|
||||
|
||||
@@ -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
|
||||
|
||||
@@ -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
|
||||
}
|
||||
|
||||
@@ -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{
|
||||
|
||||
@@ -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 {
|
||||
|
||||
@@ -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
|
||||
|
||||
@@ -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
|
||||
}
|
||||
|
||||
@@ -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
|
||||
}
|
||||
|
||||
@@ -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")
|
||||
}
|
||||
}
|
||||
|
||||
@@ -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)
|
||||
|
||||
Reference in New Issue
Block a user