diff --git a/examples/shadowsocks-dial/main.go b/examples/shadowsocks-dial/main.go new file mode 100644 index 0000000..cae3030 --- /dev/null +++ b/examples/shadowsocks-dial/main.go @@ -0,0 +1,57 @@ +package main + +import ( + "context" + "flag" + "fmt" + "io" + "log" + "net" + "net/http" + "os" + + "github.com/33TU/socks/shadowsocks" +) + +func main() { + proxyURL := flag.String( + "proxy", + os.Getenv("SS_PROXY_URL"), + `Shadowsocks proxy URL, e.g. ss://2022-blake3-aes-128-gcm:iDN+jVYAcTkUxwNICMTQRA==@127.0.0.1:8388`, + ) + + flag.Parse() + + if *proxyURL == "" { + log.Fatal("missing proxy URL; set -proxy or SS_PROXY_URL") + } + + dialer, err := shadowsocks.NewDialerFromURLString(*proxyURL, nil) + if err != nil { + log.Fatalf("failed to create dialer: %v", err) + } + + transport := &http.Transport{ + DialContext: func(ctx context.Context, network, address string) (net.Conn, error) { + return dialer.DialContext(ctx, network, address) + }, + } + + client := &http.Client{ + Transport: transport, + } + + resp, err := client.Get("https://httpbin.org/ip") + if err != nil { + log.Fatalf("request failed: %v", err) + } + defer resp.Body.Close() + + body, err := io.ReadAll(resp.Body) + if err != nil { + log.Fatalf("failed to read response body: %v", err) + } + + fmt.Printf("status: %s\n", resp.Status) + fmt.Printf("body:\n%s\n", body) +} diff --git a/go.mod b/go.mod index 46cbe21..67f4512 100644 --- a/go.mod +++ b/go.mod @@ -3,3 +3,10 @@ module github.com/33TU/socks go 1.25.1 require golang.org/x/sync v0.20.0 + +require ( + github.com/klauspost/cpuid/v2 v2.0.12 // indirect + github.com/zeebo/blake3 v0.2.4 // indirect + golang.org/x/crypto v0.50.0 // indirect + golang.org/x/sys v0.43.0 // indirect +) diff --git a/go.sum b/go.sum index 733d716..1d6683a 100644 --- a/go.sum +++ b/go.sum @@ -1,2 +1,10 @@ +github.com/klauspost/cpuid/v2 v2.0.12 h1:p9dKCg8i4gmOxtv35DvrYoWqYzQrvEVdjQ762Y0OqZE= +github.com/klauspost/cpuid/v2 v2.0.12/go.mod h1:g2LTdtYhdyuGPqyWyv7qRAmj1WBqxuObKfj5c0PQa7c= +github.com/zeebo/blake3 v0.2.4 h1:KYQPkhpRtcqh0ssGYcKLG1JYvddkEA8QwCM/yBqhaZI= +github.com/zeebo/blake3 v0.2.4/go.mod h1:7eeQ6d2iXWRGF6npfaxl2CU+xy2Fjo2gxeyZGCRUjcE= +golang.org/x/crypto v0.50.0 h1:zO47/JPrL6vsNkINmLoo/PH1gcxpls50DNogFvB5ZGI= +golang.org/x/crypto v0.50.0/go.mod h1:3muZ7vA7PBCE6xgPX7nkzzjiUq87kRItoJQM1Yo8S+Q= golang.org/x/sync v0.20.0 h1:e0PTpb7pjO8GAtTs2dQ6jYa5BWYlMuX047Dco/pItO4= golang.org/x/sync v0.20.0/go.mod h1:9xrNwdLfx4jkKbNva9FpL6vEN7evnE43NNNJQ2LF3+0= +golang.org/x/sys v0.43.0 h1:Rlag2XtaFTxp19wS8MXlJwTvoh8ArU6ezoyFsMyCTNI= +golang.org/x/sys v0.43.0/go.mod h1:4GL1E5IUh+htKOUEOaiffhrAeqysfVGipDYzABqnCmw= diff --git a/shadowsocks/addr.go b/shadowsocks/addr.go new file mode 100644 index 0000000..7e5ae37 --- /dev/null +++ b/shadowsocks/addr.go @@ -0,0 +1,189 @@ +package shadowsocks + +import ( + "encoding/binary" + "errors" + "fmt" + "net" +) + +// Common validation and decode/encode errors for Shadowsocks addresses. +var ( + ErrInvalidAddrType = errors.New("invalid address type") + ErrInvalidAddr = errors.New("invalid address") + ErrInvalidDomain = errors.New("invalid domain") + ErrShortAddr = errors.New("short address") + ErrShortAddrBuffer = errors.New("short address buffer") +) + +// Addr represents a SOCKS5-style address field used by Shadowsocks. +type Addr struct { + AddrType byte + IP net.IP + Domain string + Port uint16 +} + +// Init initializes an Addr. +func (a *Addr) Init(addrType byte, ip net.IP, domain string, port uint16) { + a.AddrType = addrType + a.IP = ip + a.Domain = domain + a.Port = port +} + +// GetHost returns the address host as either the domain or IP string. +func (a *Addr) GetHost() string { + if a.AddrType == AddrTypeDomain { + return a.Domain + } + if a.IP == nil { + return "" + } + return a.IP.String() +} + +// Addr returns the address as a combined host:port string. +func (a *Addr) Addr() string { + return net.JoinHostPort(a.GetHost(), fmt.Sprint(a.Port)) +} + +// Validate checks the correctness of the address fields. +func (a *Addr) Validate() error { + switch a.AddrType { + case AddrTypeIPv4: + if a.IP == nil || a.IP.To4() == nil { + return ErrInvalidAddr + } + case AddrTypeIPv6: + if a.IP == nil || a.IP.To16() == nil || a.IP.To4() != nil { + return ErrInvalidAddr + } + case AddrTypeDomain: + if len(a.Domain) == 0 || len(a.Domain) > 255 { + return ErrInvalidDomain + } + default: + return ErrInvalidAddrType + } + + return nil +} + +// EncodedLen returns the number of bytes required to encode the address. +func (a *Addr) EncodedLen() int { + switch a.AddrType { + case AddrTypeIPv4: + return 1 + 4 + 2 + case AddrTypeIPv6: + return 1 + 16 + 2 + case AddrTypeDomain: + return 1 + 1 + len(a.Domain) + 2 + default: + return 0 + } +} + +// Decode decodes an address from src. +// It returns the number of bytes consumed. +func (a *Addr) Decode(src []byte) (int, error) { + if len(src) < 1 { + return 0, ErrShortAddr + } + + a.AddrType = src[0] + a.IP = nil + a.Domain = "" + + switch a.AddrType { + case AddrTypeIPv4: + if len(src) < 1+4+2 { + return 0, ErrShortAddr + } + a.IP = net.IP(src[1 : 1+4]).To4() + if a.IP == nil { + return 0, ErrInvalidAddr + } + a.Port = binary.BigEndian.Uint16(src[5:7]) + return 7, nil + + case AddrTypeIPv6: + if len(src) < 1+16+2 { + return 0, ErrShortAddr + } + a.IP = net.IP(src[1 : 1+16]).To16() + if a.IP == nil || a.IP.To4() != nil { + return 0, ErrInvalidAddr + } + a.Port = binary.BigEndian.Uint16(src[17:19]) + return 19, nil + + case AddrTypeDomain: + if len(src) < 2 { + return 0, ErrShortAddr + } + n := int(src[1]) + if len(src) < 1+1+n+2 { + return 0, ErrShortAddr + } + a.Domain = string(src[2 : 2+n]) + a.Port = binary.BigEndian.Uint16(src[2+n : 2+n+2]) + if err := a.Validate(); err != nil { + return 0, err + } + return 1 + 1 + n + 2, nil + + default: + return 0, ErrInvalidAddrType + } +} + +// EncodeTo encodes the address into dst and returns the extended slice. +func (a *Addr) EncodeTo(dst []byte) ([]byte, error) { + if err := a.Validate(); err != nil { + return nil, err + } + + dst = append(dst, a.AddrType) + + switch a.AddrType { + case AddrTypeIPv4: + dst = append(dst, a.IP.To4()...) + dst = binary.BigEndian.AppendUint16(dst, a.Port) + return dst, nil + + case AddrTypeIPv6: + dst = append(dst, a.IP.To16()...) + dst = binary.BigEndian.AppendUint16(dst, a.Port) + return dst, nil + + case AddrTypeDomain: + dst = append(dst, byte(len(a.Domain))) + dst = append(dst, a.Domain...) + dst = binary.BigEndian.AppendUint16(dst, a.Port) + return dst, nil + + default: + return nil, ErrInvalidAddrType + } +} + +// String returns a human-readable representation of the address. +func (a *Addr) String() string { + var atype string + switch a.AddrType { + case AddrTypeIPv4: + atype = "IPv4" + case AddrTypeDomain: + atype = "DOMAIN" + case AddrTypeIPv6: + atype = "IPv6" + default: + atype = fmt.Sprintf("0x%02X", a.AddrType) + } + + return fmt.Sprintf( + "Addr{AddrType=%s, Host=%s, Port=%d}", + atype, a.GetHost(), a.Port, + ) +} diff --git a/shadowsocks/addr_test.go b/shadowsocks/addr_test.go new file mode 100644 index 0000000..5764979 --- /dev/null +++ b/shadowsocks/addr_test.go @@ -0,0 +1,462 @@ +package shadowsocks_test + +import ( + "errors" + "net" + "testing" + + "github.com/33TU/socks/shadowsocks" +) + +func TestAddr_Init_Validate(t *testing.T) { + tests := []struct { + name string + addr shadowsocks.Addr + wantErr error + }{ + { + name: "valid ipv4", + addr: func() shadowsocks.Addr { + var a shadowsocks.Addr + a.Init(shadowsocks.AddrTypeIPv4, net.IPv4(127, 0, 0, 1), "", 1080) + return a + }(), + }, + { + name: "valid domain", + addr: func() shadowsocks.Addr { + var a shadowsocks.Addr + a.Init(shadowsocks.AddrTypeDomain, nil, "example.com", 443) + return a + }(), + }, + { + name: "valid ipv6", + addr: func() shadowsocks.Addr { + var a shadowsocks.Addr + a.Init(shadowsocks.AddrTypeIPv6, net.ParseIP("::1"), "", 5353) + return a + }(), + }, + { + name: "invalid addr type", + addr: func() shadowsocks.Addr { + var a shadowsocks.Addr + a.Init(0x99, net.IPv4(127, 0, 0, 1), "", 1080) + return a + }(), + wantErr: shadowsocks.ErrInvalidAddrType, + }, + { + name: "invalid ipv4 ip", + addr: func() shadowsocks.Addr { + var a shadowsocks.Addr + a.Init(shadowsocks.AddrTypeIPv4, nil, "", 1080) + return a + }(), + wantErr: shadowsocks.ErrInvalidAddr, + }, + { + name: "invalid ipv6 ip", + addr: func() shadowsocks.Addr { + var a shadowsocks.Addr + a.Init(shadowsocks.AddrTypeIPv6, net.IPv4(127, 0, 0, 1), "", 1080) + return a + }(), + wantErr: shadowsocks.ErrInvalidAddr, + }, + { + name: "missing domain", + addr: func() shadowsocks.Addr { + var a shadowsocks.Addr + a.Init(shadowsocks.AddrTypeDomain, nil, "", 1080) + return a + }(), + wantErr: shadowsocks.ErrInvalidDomain, + }, + { + name: "domain too long", + addr: func() shadowsocks.Addr { + var a shadowsocks.Addr + a.Init(shadowsocks.AddrTypeDomain, nil, string(make([]byte, 256)), 1080) + return a + }(), + wantErr: shadowsocks.ErrInvalidDomain, + }, + } + + for _, tt := range tests { + t.Run(tt.name, func(t *testing.T) { + err := tt.addr.Validate() + if !errors.Is(err, tt.wantErr) { + t.Errorf("Validate() error = %v, wantErr = %v", err, tt.wantErr) + } + }) + } +} + +func TestAddr_GetHost(t *testing.T) { + tests := []struct { + name string + addr shadowsocks.Addr + want string + }{ + { + name: "domain", + addr: shadowsocks.Addr{ + AddrType: shadowsocks.AddrTypeDomain, + Domain: "example.com", + Port: 443, + }, + want: "example.com", + }, + { + name: "ipv4", + addr: shadowsocks.Addr{ + AddrType: shadowsocks.AddrTypeIPv4, + IP: net.IPv4(127, 0, 0, 1), + Port: 1080, + }, + want: "127.0.0.1", + }, + { + name: "nil ip", + addr: shadowsocks.Addr{ + AddrType: shadowsocks.AddrTypeIPv4, + IP: nil, + Port: 1080, + }, + want: "", + }, + } + + for _, tt := range tests { + t.Run(tt.name, func(t *testing.T) { + got := tt.addr.GetHost() + if got != tt.want { + t.Fatalf("GetHost() = %q, want %q", got, tt.want) + } + }) + } +} + +func TestAddr_Addr(t *testing.T) { + tests := []struct { + name string + addr shadowsocks.Addr + want string + }{ + { + name: "domain", + addr: shadowsocks.Addr{ + AddrType: shadowsocks.AddrTypeDomain, + Domain: "example.com", + Port: 443, + }, + want: "example.com:443", + }, + { + name: "ipv4", + addr: shadowsocks.Addr{ + AddrType: shadowsocks.AddrTypeIPv4, + IP: net.IPv4(127, 0, 0, 1), + Port: 1080, + }, + want: "127.0.0.1:1080", + }, + { + name: "ipv6", + addr: shadowsocks.Addr{ + AddrType: shadowsocks.AddrTypeIPv6, + IP: net.ParseIP("::1"), + Port: 53, + }, + want: "[::1]:53", + }, + } + + for _, tt := range tests { + t.Run(tt.name, func(t *testing.T) { + got := tt.addr.Addr() + if got != tt.want { + t.Fatalf("Addr() = %q, want %q", got, tt.want) + } + }) + } +} + +func TestAddr_EncodedLen(t *testing.T) { + tests := []struct { + name string + addr shadowsocks.Addr + want int + }{ + { + name: "ipv4", + addr: shadowsocks.Addr{ + AddrType: shadowsocks.AddrTypeIPv4, + IP: net.IPv4(127, 0, 0, 1), + Port: 1080, + }, + want: 7, + }, + { + name: "domain", + addr: shadowsocks.Addr{ + AddrType: shadowsocks.AddrTypeDomain, + Domain: "example.com", + Port: 443, + }, + want: 1 + 1 + len("example.com") + 2, + }, + { + name: "ipv6", + addr: shadowsocks.Addr{ + AddrType: shadowsocks.AddrTypeIPv6, + IP: net.ParseIP("::1"), + Port: 53, + }, + want: 19, + }, + { + name: "invalid", + addr: shadowsocks.Addr{ + AddrType: 0x99, + Port: 1, + }, + want: 0, + }, + } + + for _, tt := range tests { + t.Run(tt.name, func(t *testing.T) { + got := tt.addr.EncodedLen() + if got != tt.want { + t.Errorf("EncodedLen() = %d, want %d", got, tt.want) + } + }) + } +} + +func TestAddr_EncodeTo_Decode_RoundTrip(t *testing.T) { + tests := []struct { + name string + addr shadowsocks.Addr + }{ + { + name: "ipv4", + addr: shadowsocks.Addr{ + AddrType: shadowsocks.AddrTypeIPv4, + IP: net.IPv4(127, 0, 0, 1), + Port: 1080, + }, + }, + { + name: "domain", + addr: shadowsocks.Addr{ + AddrType: shadowsocks.AddrTypeDomain, + Domain: "example.com", + Port: 443, + }, + }, + { + name: "ipv6", + addr: shadowsocks.Addr{ + AddrType: shadowsocks.AddrTypeIPv6, + IP: net.ParseIP("2001:db8::1"), + Port: 5353, + }, + }, + } + + for _, tt := range tests { + t.Run(tt.name, func(t *testing.T) { + buf := make([]byte, tt.addr.EncodedLen()) + + bw, err := tt.addr.EncodeTo(buf[:0]) + if err != nil { + t.Fatalf("EncodeTo() failed: %v", err) + } + if len(bw) != len(buf) { + t.Fatalf("EncodeTo() wrote %d bytes, want %d", len(bw), len(buf)) + } + + var got shadowsocks.Addr + nr, err := got.Decode(buf) + if err != nil { + t.Fatalf("Decode() failed: %v", err) + } + if nr != len(buf) { + t.Fatalf("Decode() read %d bytes, want %d", nr, len(buf)) + } + + if got.AddrType != tt.addr.AddrType { + t.Fatalf("AddrType = %v, want %v", got.AddrType, tt.addr.AddrType) + } + if got.Port != tt.addr.Port { + t.Fatalf("Port = %d, want %d", got.Port, tt.addr.Port) + } + if got.Domain != tt.addr.Domain { + t.Fatalf("Domain = %q, want %q", got.Domain, tt.addr.Domain) + } + + switch tt.addr.AddrType { + case shadowsocks.AddrTypeIPv4: + if !got.IP.Equal(tt.addr.IP.To4()) { + t.Fatalf("IP = %v, want %v", got.IP, tt.addr.IP.To4()) + } + case shadowsocks.AddrTypeIPv6: + if !got.IP.Equal(tt.addr.IP.To16()) { + t.Fatalf("IP = %v, want %v", got.IP, tt.addr.IP.To16()) + } + } + }) + } +} + +func TestAddr_EncodeTo_Invalid(t *testing.T) { + tests := []struct { + name string + addr shadowsocks.Addr + bufLen int + wantErr error + }{ + { + name: "invalid addr type", + addr: shadowsocks.Addr{ + AddrType: 0x99, + IP: net.IPv4(127, 0, 0, 1), + Port: 1080, + }, + bufLen: 32, + wantErr: shadowsocks.ErrInvalidAddrType, + }, + { + name: "invalid domain", + addr: shadowsocks.Addr{ + AddrType: shadowsocks.AddrTypeDomain, + Domain: "", + Port: 80, + }, + bufLen: 32, + wantErr: shadowsocks.ErrInvalidDomain, + }, + } + + for _, tt := range tests { + t.Run(tt.name, func(t *testing.T) { + buf := make([]byte, tt.bufLen) + _, err := tt.addr.EncodeTo(buf[:0]) + if !errors.Is(err, tt.wantErr) { + t.Fatalf("EncodeTo() error = %v, wantErr = %v", err, tt.wantErr) + } + }) + } +} + +func TestAddr_Decode_Invalid(t *testing.T) { + tests := []struct { + name string + src []byte + wantErr error + }{ + { + name: "empty", + src: nil, + wantErr: shadowsocks.ErrShortAddr, + }, + { + name: "invalid addr type", + src: []byte{0x99}, + wantErr: shadowsocks.ErrInvalidAddrType, + }, + { + name: "short ipv4", + src: []byte{shadowsocks.AddrTypeIPv4, 127, 0, 0}, + wantErr: shadowsocks.ErrShortAddr, + }, + { + name: "short domain length byte missing payload", + src: []byte{shadowsocks.AddrTypeDomain, 5, 'a', 'b'}, + wantErr: shadowsocks.ErrShortAddr, + }, + { + name: "short ipv6", + src: []byte{shadowsocks.AddrTypeIPv6, 0, 1, 2}, + wantErr: shadowsocks.ErrShortAddr, + }, + { + name: "empty domain", + src: []byte{ + shadowsocks.AddrTypeDomain, + 0, + 0, 80, + }, + wantErr: shadowsocks.ErrInvalidDomain, + }, + } + + for _, tt := range tests { + t.Run(tt.name, func(t *testing.T) { + var a shadowsocks.Addr + _, err := a.Decode(tt.src) + if !errors.Is(err, tt.wantErr) { + t.Fatalf("Decode() error = %v, wantErr = %v", err, tt.wantErr) + } + }) + } +} + +func TestAddr_String(t *testing.T) { + tests := []struct { + name string + addr shadowsocks.Addr + want string + }{ + { + name: "ipv4", + addr: shadowsocks.Addr{ + AddrType: shadowsocks.AddrTypeIPv4, + IP: net.IPv4(127, 0, 0, 1), + Port: 1080, + }, + want: "Addr{AddrType=IPv4, Host=127.0.0.1, Port=1080}", + }, + { + name: "domain", + addr: shadowsocks.Addr{ + AddrType: shadowsocks.AddrTypeDomain, + Domain: "example.com", + Port: 443, + }, + want: "Addr{AddrType=DOMAIN, Host=example.com, Port=443}", + }, + { + name: "ipv6", + addr: shadowsocks.Addr{ + AddrType: shadowsocks.AddrTypeIPv6, + IP: net.ParseIP("::1"), + Port: 53, + }, + want: "Addr{AddrType=IPv6, Host=::1, Port=53}", + }, + { + name: "unknown", + addr: shadowsocks.Addr{ + AddrType: 0x99, + Domain: "x", + Port: 1, + }, + want: "Addr{AddrType=0x99, Host=, Port=1}", + }, + } + + for _, tt := range tests { + t.Run(tt.name, func(t *testing.T) { + got := tt.addr.String() + if got != tt.want { + t.Fatalf("String() = %q, want %q", got, tt.want) + } + }) + } +} diff --git a/shadowsocks/config.go b/shadowsocks/config.go new file mode 100644 index 0000000..eeeed1a --- /dev/null +++ b/shadowsocks/config.go @@ -0,0 +1,51 @@ +package shadowsocks + +import ( + "encoding/base64" + "fmt" +) + +// Config holds the configuration for a Shadowsocks proxy. +type Config struct { + Method string + PSK string + Plugin string + Tag string +} + +// Validate checks if the Config is valid and returns an error if not. +func (c *Config) Validate() error { + if c == nil { + return fmt.Errorf("nil config") + } + + switch c.Method { + case Method2022Blake3AES128GCM, + Method2022Blake3AES256GCM, + Method2022Blake3ChaCha20Poly1305: + default: + return fmt.Errorf("invalid method: %s", c.Method) + } + + if c.PSK == "" { + return fmt.Errorf("missing PSK") + } + + rawKey, err := base64.StdEncoding.DecodeString(c.PSK) + if err != nil { + return fmt.Errorf("invalid PSK: must be standard base64: %w", err) + } + + switch c.Method { + case Method2022Blake3AES128GCM: + if len(rawKey) != 16 { + return fmt.Errorf("invalid PSK length for %s: got %d, want 16", c.Method, len(rawKey)) + } + case Method2022Blake3AES256GCM, Method2022Blake3ChaCha20Poly1305: + if len(rawKey) != 32 { + return fmt.Errorf("invalid PSK length for %s: got %d, want 32", c.Method, len(rawKey)) + } + } + + return nil +} diff --git a/shadowsocks/config_test.go b/shadowsocks/config_test.go new file mode 100644 index 0000000..7e962a0 --- /dev/null +++ b/shadowsocks/config_test.go @@ -0,0 +1,120 @@ +package shadowsocks_test + +import ( + "encoding/base64" + "strings" + "testing" + + "github.com/33TU/socks/shadowsocks" +) + +func TestConfig_Validate(t *testing.T) { + t.Parallel() + + psk16 := base64.StdEncoding.EncodeToString(make([]byte, 16)) + psk32 := base64.StdEncoding.EncodeToString(make([]byte, 32)) + psk15 := base64.StdEncoding.EncodeToString(make([]byte, 15)) + psk31 := base64.StdEncoding.EncodeToString(make([]byte, 31)) + + tests := []struct { + name string + cfg *shadowsocks.Config + wantErr string + }{ + { + name: "nil config", + cfg: nil, + wantErr: "nil config", + }, + { + name: "valid aes128", + cfg: &shadowsocks.Config{ + Method: shadowsocks.Method2022Blake3AES128GCM, + PSK: psk16, + }, + }, + { + name: "valid aes256", + cfg: &shadowsocks.Config{ + Method: shadowsocks.Method2022Blake3AES256GCM, + PSK: psk32, + }, + }, + { + name: "valid chacha20", + cfg: &shadowsocks.Config{ + Method: shadowsocks.Method2022Blake3ChaCha20Poly1305, + PSK: psk32, + }, + }, + { + name: "invalid method", + cfg: &shadowsocks.Config{ + Method: "invalid-method", + PSK: psk32, + }, + wantErr: "invalid method", + }, + { + name: "missing psk", + cfg: &shadowsocks.Config{ + Method: shadowsocks.Method2022Blake3AES128GCM, + }, + wantErr: "missing PSK", + }, + { + name: "invalid base64", + cfg: &shadowsocks.Config{ + Method: shadowsocks.Method2022Blake3AES128GCM, + PSK: "!!!not-base64!!!", + }, + wantErr: "invalid PSK: must be standard base64", + }, + { + name: "invalid aes128 key length", + cfg: &shadowsocks.Config{ + Method: shadowsocks.Method2022Blake3AES128GCM, + PSK: psk15, + }, + wantErr: "invalid PSK length", + }, + { + name: "invalid aes256 key length", + cfg: &shadowsocks.Config{ + Method: shadowsocks.Method2022Blake3AES256GCM, + PSK: psk31, + }, + wantErr: "invalid PSK length", + }, + { + name: "invalid chacha20 key length", + cfg: &shadowsocks.Config{ + Method: shadowsocks.Method2022Blake3ChaCha20Poly1305, + PSK: psk31, + }, + wantErr: "invalid PSK length", + }, + } + + for _, tt := range tests { + tt := tt + t.Run(tt.name, func(t *testing.T) { + t.Parallel() + + err := tt.cfg.Validate() + if tt.wantErr == "" { + if err != nil { + t.Fatalf("Validate() error = %v", err) + } + return + } + + if err == nil { + t.Fatalf("expected error containing %q, got nil", tt.wantErr) + } + if !strings.Contains(err.Error(), tt.wantErr) { + t.Fatalf("expected error containing %q, got %q", tt.wantErr, err.Error()) + } + }) + } +} diff --git a/shadowsocks/consts.go b/shadowsocks/consts.go new file mode 100644 index 0000000..a01c107 --- /dev/null +++ b/shadowsocks/consts.go @@ -0,0 +1,15 @@ +package shadowsocks + +// Encryption method constants for Shadowsocks AEAD-2022. +const ( + Method2022Blake3AES128GCM = "2022-blake3-aes-128-gcm" + Method2022Blake3AES256GCM = "2022-blake3-aes-256-gcm" + Method2022Blake3ChaCha20Poly1305 = "2022-blake3-chacha20-poly1305" +) + +// SOCKS5-style address types used inside Shadowsocks headers. +const ( + AddrTypeIPv4 = 0x01 + AddrTypeDomain = 0x03 + AddrTypeIPv6 = 0x04 +) diff --git a/shadowsocks/dialer.go b/shadowsocks/dialer.go new file mode 100644 index 0000000..f105fe1 --- /dev/null +++ b/shadowsocks/dialer.go @@ -0,0 +1,249 @@ +package shadowsocks + +import ( + "context" + "fmt" + "net" + "net/url" + "time" + + "github.com/33TU/socks/internal" + socksnet "github.com/33TU/socks/net" +) + +// Dialer implements a Shadowsocks proxy dialer. +type Dialer struct { + ProxyAddr string + Config *Config + Dialer socksnet.Dialer +} + +// NewDialer creates a new Shadowsocks dialer instance. +func NewDialer(proxyAddr string, cfg *Config, dialer socksnet.Dialer) *Dialer { + if dialer == nil { + dialer = socksnet.DefaultDialer + } + + return &Dialer{ + ProxyAddr: proxyAddr, + Config: cfg, + Dialer: dialer, + } +} + +// NewDialerFromURL creates a new Dialer from a URL of the form +// ss://method:psk@host:port[/?plugin=...][#tag] +// +// For AEAD-2022, userinfo must be plain method:psk and not legacy base64-wrapped userinfo. +func NewDialerFromURL(u *url.URL, dialer socksnet.Dialer) (*Dialer, error) { + if u == nil { + return nil, fmt.Errorf("nil proxy URL") + } + + switch u.Scheme { + case "ss": + default: + return nil, fmt.Errorf("invalid scheme: %s", u.Scheme) + } + + host := u.Hostname() + if host == "" { + return nil, fmt.Errorf("missing host in proxy URL") + } + + port := u.Port() + if port == "" { + return nil, fmt.Errorf("missing port in proxy URL") + } + + proxyAddr := net.JoinHostPort(host, port) + + if u.User == nil { + return nil, fmt.Errorf("missing method/psk in proxy URL") + } + + method := u.User.Username() + if method == "" { + return nil, fmt.Errorf("missing method in proxy URL") + } + + psk, hasPassword := u.User.Password() + if !hasPassword { + return nil, fmt.Errorf("missing PSK in proxy URL") + } + + cfg := &Config{ + Method: method, + PSK: psk, + Plugin: u.Query().Get("plugin"), + Tag: u.Fragment, + } + if err := cfg.Validate(); err != nil { + return nil, fmt.Errorf("invalid proxy URL: %w", err) + } + + return NewDialer(proxyAddr, cfg, dialer), nil +} + +// NewDialerFromURLString creates a new Dialer from a URL string of the form +// ss://method:psk@host:port[/?plugin=...][#tag] +// +// For AEAD-2022, userinfo must be plain method:psk and not legacy base64-wrapped userinfo. +func NewDialerFromURLString(rawURL string, dialer socksnet.Dialer) (*Dialer, error) { + u, err := url.Parse(rawURL) + if err != nil { + return nil, fmt.Errorf("invalid proxy URL: %w", err) + } + return NewDialerFromURL(u, dialer) +} + +// ProxyAddress returns the configured Shadowsocks proxy address. +func (d *Dialer) ProxyAddress() string { + return d.ProxyAddr +} + +// DialContext establishes a TCP connection via the Shadowsocks proxy. +func (d *Dialer) DialContext(ctx context.Context, network, address string) (net.Conn, error) { + conn, err := d.dialProxy(ctx, network) + if err != nil { + return nil, err + } + + return d.DialConnContext(ctx, conn, network, address) +} + +// Dial establishes a TCP connection via the Shadowsocks proxy using background context. +func (d *Dialer) Dial(network, address string) (net.Conn, error) { + return d.DialContext(context.Background(), network, address) +} + +// DialConnContext upgrades an existing connection into a Shadowsocks TCP stream. +func (d *Dialer) DialConnContext(ctx context.Context, conn net.Conn, network, address string) (net.Conn, error) { + if d == nil { + conn.Close() + return nil, fmt.Errorf("nil shadowsocks dialer") + } + if d.Config == nil { + conn.Close() + return nil, fmt.Errorf("missing shadowsocks config") + } + if err := d.Config.Validate(); err != nil { + conn.Close() + return nil, err + } + + method, err := ParseMethod(d.Config.Method) + if err != nil { + conn.Close() + return nil, err + } + + psk, err := DecodePSKTo(nil, method, d.Config.PSK) + if err != nil { + conn.Close() + return nil, err + } + + target, err := parseTargetAddr(address) + if err != nil { + conn.Close() + return nil, err + } + + // cancellation and deadline handling + cleanup := bindConnToContext(ctx, conn) + defer cleanup() + + paddingLen, err := RandomInt(1, 64) // make configurable later + if err != nil { + conn.Close() + return nil, err + } + + padding := internal.GetBytes(paddingLen) + defer internal.PutBytes(padding) + + if err := FillRandomBytes(padding); err != nil { + conn.Close() + return nil, err + } + + ssConn, err := NewClientTCPConn(conn, method, psk, target, padding, nil) + if err != nil { + conn.Close() + return nil, err + } + return ssConn, nil +} + +// DialConn upgrades an existing connection using background context. +func (d *Dialer) DialConn(conn net.Conn, network, address string) (net.Conn, error) { + return d.DialConnContext(context.Background(), conn, network, address) +} + +// dialProxy connects to the Shadowsocks proxy server. +func (d *Dialer) dialProxy(ctx context.Context, network string) (net.Conn, error) { + dialer := d.Dialer + if dialer == nil { + dialer = socksnet.DefaultDialer + } + return dialer.DialContext(ctx, network, d.ProxyAddr) +} + +func parseTargetAddr(address string) (Addr, error) { + host, portStr, err := net.SplitHostPort(address) + if err != nil { + return Addr{}, err + } + + port, err := net.DefaultResolver.LookupPort(context.Background(), "tcp", portStr) + if err != nil { + return Addr{}, err + } + + ip := net.ParseIP(host) + switch { + case ip != nil && ip.To4() != nil: + return Addr{ + AddrType: AddrTypeIPv4, + IP: ip.To4(), + Port: uint16(port), + }, nil + + case ip != nil && ip.To16() != nil: + return Addr{ + AddrType: AddrTypeIPv6, + IP: ip.To16(), + Port: uint16(port), + }, nil + + default: + return Addr{ + AddrType: AddrTypeDomain, + Domain: host, + Port: uint16(port), + }, nil + } +} + +// bindConnToContext sets connection deadlines based on context and ensures cleanup on cancellation. +func bindConnToContext(ctx context.Context, conn net.Conn) (cleanup func()) { + if deadline, ok := ctx.Deadline(); ok { + _ = conn.SetDeadline(deadline) + } + + exitCh := make(chan struct{}) + + go func() { + select { + case <-ctx.Done(): + _ = conn.Close() + case <-exitCh: + } + }() + + return func() { + close(exitCh) + _ = conn.SetDeadline(time.Time{}) + } +} diff --git a/shadowsocks/dialer_test.go b/shadowsocks/dialer_test.go new file mode 100644 index 0000000..5cfa6d5 --- /dev/null +++ b/shadowsocks/dialer_test.go @@ -0,0 +1,508 @@ +package shadowsocks_test + +import ( + "bytes" + "context" + "encoding/base64" + "io" + "net" + "net/url" + "strings" + "testing" + "time" + + socksnet "github.com/33TU/socks/net" + "github.com/33TU/socks/shadowsocks" +) + +func mustBase64Key(n int) string { + return base64.StdEncoding.EncodeToString(bytes.Repeat([]byte{0x42}, n)) +} + +func startMockShadowsocksServer(t *testing.T, handle func(net.Conn)) (string, func()) { + t.Helper() + + ln, err := net.Listen("tcp", "127.0.0.1:0") + if err != nil { + t.Fatalf("listen: %v", err) + } + + go func() { + for { + conn, err := ln.Accept() + if err != nil { + return + } + go handle(conn) + } + }() + + return ln.Addr().String(), func() { _ = ln.Close() } +} + +func decodePSKForMethod(t *testing.T, method shadowsocks.Method, psk string) []byte { + t.Helper() + + raw, err := shadowsocks.DecodePSKTo(nil, method, psk) + if err != nil { + t.Fatalf("DecodePSKTo() error = %v", err) + } + return raw +} + +func TestNewDialer(t *testing.T) { + t.Parallel() + + cfg := &shadowsocks.Config{ + Method: shadowsocks.Method2022Blake3AES256GCM, + PSK: mustBase64Key(32), + } + + d := shadowsocks.NewDialer("127.0.0.1:8388", cfg, nil) + if d == nil { + t.Fatal("NewDialer() returned nil") + } + if d.ProxyAddr != "127.0.0.1:8388" { + t.Fatalf("ProxyAddr = %q, want %q", d.ProxyAddr, "127.0.0.1:8388") + } + if d.Config != cfg { + t.Fatal("Config pointer mismatch") + } + if d.Dialer == nil { + t.Fatal("Dialer is nil, want default dialer") + } +} + +func TestDialer_ProxyAddress(t *testing.T) { + t.Parallel() + + d := &shadowsocks.Dialer{ProxyAddr: "127.0.0.1:8388"} + if got := d.ProxyAddress(); got != "127.0.0.1:8388" { + t.Fatalf("ProxyAddress() = %q, want %q", got, "127.0.0.1:8388") + } +} + +func TestNewDialerFromURL(t *testing.T) { + t.Parallel() + + tests := []struct { + name string + rawURL string + wantAddr string + wantMethod string + wantPSK string + wantPlugin string + wantTag string + wantErr string + }{ + { + name: "valid aes-256-gcm", + rawURL: "ss://2022-blake3-aes-256-gcm:" + url.QueryEscape(mustBase64Key(32)) + "@127.0.0.1:8388", + wantAddr: "127.0.0.1:8388", + wantMethod: shadowsocks.Method2022Blake3AES256GCM, + wantPSK: mustBase64Key(32), + }, + { + name: "valid aes-128-gcm with plugin and tag", + rawURL: "ss://2022-blake3-aes-128-gcm:" + url.QueryEscape(mustBase64Key(16)) + "@proxy.example.com:443/?plugin=" + url.QueryEscape("v2ray-plugin;server;host=example.com") + "#edge", + wantAddr: "proxy.example.com:443", + wantMethod: shadowsocks.Method2022Blake3AES128GCM, + wantPSK: mustBase64Key(16), + wantPlugin: "v2ray-plugin;server;host=example.com", + wantTag: "edge", + }, + { + name: "valid chacha20-poly1305", + rawURL: "ss://2022-blake3-chacha20-poly1305:" + url.QueryEscape(mustBase64Key(32)) + "@example.com:8443#demo", + wantAddr: "example.com:8443", + wantMethod: shadowsocks.Method2022Blake3ChaCha20Poly1305, + wantPSK: mustBase64Key(32), + wantTag: "demo", + }, + { + name: "nil url", + wantErr: "nil proxy URL", + }, + { + name: "invalid scheme", + rawURL: "http://127.0.0.1:8388", + wantErr: "invalid scheme", + }, + { + name: "missing host", + rawURL: "ss://2022-blake3-aes-256-gcm:" + url.QueryEscape(mustBase64Key(32)) + "@:8388", + wantErr: "missing host", + }, + { + name: "missing port", + rawURL: "ss://2022-blake3-aes-256-gcm:" + url.QueryEscape(mustBase64Key(32)) + "@127.0.0.1", + wantErr: "missing port", + }, + { + name: "missing userinfo", + rawURL: "ss://127.0.0.1:8388", + wantErr: "missing method/psk", + }, + { + name: "missing method", + rawURL: "ss://:" + url.QueryEscape(mustBase64Key(32)) + "@127.0.0.1:8388", + wantErr: "missing method", + }, + { + name: "missing psk", + rawURL: "ss://2022-blake3-aes-256-gcm@127.0.0.1:8388", + wantErr: "missing PSK", + }, + { + name: "invalid method", + rawURL: "ss://aes-256-gcm:" + url.QueryEscape(mustBase64Key(32)) + "@127.0.0.1:8388", + wantErr: "invalid proxy URL", + }, + { + name: "invalid psk length for aes-128", + rawURL: "ss://2022-blake3-aes-128-gcm:" + url.QueryEscape(mustBase64Key(32)) + "@127.0.0.1:8388", + wantErr: "invalid proxy URL", + }, + { + name: "invalid psk base64", + rawURL: "ss://2022-blake3-aes-256-gcm:not-base64!!!@127.0.0.1:8388", + wantErr: "invalid proxy URL", + }, + { + name: "ipv6 host", + rawURL: "ss://2022-blake3-aes-256-gcm:" + url.QueryEscape(mustBase64Key(32)) + "@[2001:db8::1]:8388", + wantAddr: "[2001:db8::1]:8388", + wantMethod: shadowsocks.Method2022Blake3AES256GCM, + wantPSK: mustBase64Key(32), + }, + } + + for _, tt := range tests { + tt := tt + t.Run(tt.name, func(t *testing.T) { + t.Parallel() + + if tt.wantErr == "nil proxy URL" { + d, err := shadowsocks.NewDialerFromURL(nil, nil) + if d != nil { + t.Fatal("expected nil dialer") + } + if err == nil || !strings.Contains(err.Error(), tt.wantErr) { + t.Fatalf("expected error containing %q, got %v", tt.wantErr, err) + } + return + } + + u, err := url.Parse(tt.rawURL) + if err != nil { + t.Fatalf("url.Parse(%q): %v", tt.rawURL, err) + } + + d, err := shadowsocks.NewDialerFromURL(u, nil) + if tt.wantErr != "" { + if err == nil { + t.Fatalf("expected error containing %q, got nil", tt.wantErr) + } + if !strings.Contains(err.Error(), tt.wantErr) { + t.Fatalf("expected error containing %q, got %q", tt.wantErr, err.Error()) + } + return + } + if err != nil { + t.Fatalf("NewDialerFromURL() error = %v", err) + } + if d == nil { + t.Fatal("NewDialerFromURL() returned nil dialer") + } + if d.ProxyAddr != tt.wantAddr { + t.Fatalf("ProxyAddr = %q, want %q", d.ProxyAddr, tt.wantAddr) + } + if d.Config == nil { + t.Fatal("Config is nil") + } + if d.Config.Method != tt.wantMethod { + t.Fatalf("Config.Method = %q, want %q", d.Config.Method, tt.wantMethod) + } + if d.Config.PSK != tt.wantPSK { + t.Fatalf("Config.PSK = %q, want %q", d.Config.PSK, tt.wantPSK) + } + if d.Config.Plugin != tt.wantPlugin { + t.Fatalf("Config.Plugin = %q, want %q", d.Config.Plugin, tt.wantPlugin) + } + if d.Config.Tag != tt.wantTag { + t.Fatalf("Config.Tag = %q, want %q", d.Config.Tag, tt.wantTag) + } + if d.Dialer == nil { + t.Fatal("Dialer is nil, want default dialer") + } + }) + } +} + +func TestNewDialerFromURLString(t *testing.T) { + t.Parallel() + + rawURL := "ss://2022-blake3-aes-256-gcm:" + url.QueryEscape(mustBase64Key(32)) + "@127.0.0.1:8388/?plugin=" + url.QueryEscape("simple-obfs") + "#local" + + d, err := shadowsocks.NewDialerFromURLString(rawURL, socksnet.DefaultDialer) + if err != nil { + t.Fatalf("NewDialerFromURLString() error = %v", err) + } + if d == nil { + t.Fatal("NewDialerFromURLString() returned nil dialer") + } + if d.ProxyAddr != "127.0.0.1:8388" { + t.Fatalf("ProxyAddr = %q, want %q", d.ProxyAddr, "127.0.0.1:8388") + } + if d.Config == nil { + t.Fatal("Config is nil") + } + if d.Config.Method != shadowsocks.Method2022Blake3AES256GCM { + t.Fatalf("Config.Method = %q, want %q", d.Config.Method, shadowsocks.Method2022Blake3AES256GCM) + } + if d.Config.Plugin != "simple-obfs" { + t.Fatalf("Config.Plugin = %q, want %q", d.Config.Plugin, "simple-obfs") + } + if d.Config.Tag != "local" { + t.Fatalf("Config.Tag = %q, want %q", d.Config.Tag, "local") + } +} + +func TestNewDialerFromURLString_Invalid(t *testing.T) { + t.Parallel() + + _, err := shadowsocks.NewDialerFromURLString("://bad url", nil) + if err == nil { + t.Fatal("expected error, got nil") + } + if !strings.Contains(err.Error(), "invalid proxy URL") { + t.Fatalf("unexpected error: %v", err) + } +} + +func TestDialer_DialConnContext_Errors(t *testing.T) { + t.Parallel() + + t.Run("nil dialer", func(t *testing.T) { + t.Parallel() + + c1, c2 := net.Pipe() + defer c2.Close() + + var d *shadowsocks.Dialer + _, err := d.DialConnContext(context.Background(), c1, "tcp", "example.com:443") + if err == nil { + t.Fatal("expected error, got nil") + } + if !strings.Contains(err.Error(), "nil shadowsocks dialer") { + t.Fatalf("unexpected error: %v", err) + } + }) + + t.Run("missing config", func(t *testing.T) { + t.Parallel() + + c1, c2 := net.Pipe() + defer c2.Close() + + d := &shadowsocks.Dialer{} + _, err := d.DialConnContext(context.Background(), c1, "tcp", "example.com:443") + if err == nil { + t.Fatal("expected error, got nil") + } + if !strings.Contains(err.Error(), "missing shadowsocks config") { + t.Fatalf("unexpected error: %v", err) + } + }) + + t.Run("invalid config", func(t *testing.T) { + t.Parallel() + + c1, c2 := net.Pipe() + defer c2.Close() + + d := &shadowsocks.Dialer{ + Config: &shadowsocks.Config{ + Method: shadowsocks.Method2022Blake3AES128GCM, + PSK: mustBase64Key(32), + }, + } + _, err := d.DialConnContext(context.Background(), c1, "tcp", "example.com:443") + if err == nil { + t.Fatal("expected error, got nil") + } + }) + + t.Run("invalid target address", func(t *testing.T) { + t.Parallel() + + c1, c2 := net.Pipe() + defer c2.Close() + + d := &shadowsocks.Dialer{ + Config: &shadowsocks.Config{ + Method: shadowsocks.Method2022Blake3AES128GCM, + PSK: mustBase64Key(16), + }, + } + _, err := d.DialConnContext(context.Background(), c1, "tcp", "not-a-hostport") + if err == nil { + t.Fatal("expected error, got nil") + } + }) +} + +func TestDialer_DialContext_Success(t *testing.T) { + t.Parallel() + + methodName := shadowsocks.Method2022Blake3AES128GCM + pskB64 := mustBase64Key(16) + + method, err := shadowsocks.ParseMethod(methodName) + if err != nil { + t.Fatalf("ParseMethod() error = %v", err) + } + psk := decodePSKForMethod(t, method, pskB64) + + proxyAddr, stop := startMockShadowsocksServer(t, func(c net.Conn) { + defer c.Close() + + reqStart, _, err := shadowsocks.ReadTCPRequestStart(c, method, psk) + if err != nil { + t.Errorf("server: ReadTCPRequestStart() error = %v", err) + return + } + + if reqStart.Header.Target.AddrType != shadowsocks.AddrTypeDomain { + t.Errorf("server: target type = %v, want domain", reqStart.Header.Target.AddrType) + return + } + if reqStart.Header.Target.Domain != "example.com" { + t.Errorf("server: target domain = %q, want %q", reqStart.Header.Target.Domain, "example.com") + return + } + if reqStart.Header.Target.Port != 443 { + t.Errorf("server: target port = %d, want %d", reqStart.Header.Target.Port, 443) + return + } + + var reader shadowsocks.TCPChunkReader + if err := reader.Init(reqStart.Cipher, c); err != nil { + t.Errorf("server: reader Init() error = %v", err) + return + } + + payload, _, err := reader.ReadChunkTo(nil) + if err != nil { + t.Errorf("server: ReadChunkTo() error = %v", err) + return + } + if string(payload) != "ping" { + t.Errorf("server: payload = %q, want %q", payload, "ping") + return + } + + responseSalt := bytes.Repeat([]byte{0x33}, method.SaltSize) + initialPayload := []byte("pong") + + _, _, err = shadowsocks.WriteTCPResponseStart( + c, + method, + psk, + responseSalt, + time.Now(), + reqStart.Salt, + initialPayload, + ) + if err != nil { + t.Errorf("server: WriteTCPResponseStart() error = %v", err) + return + } + }) + defer stop() + + d := shadowsocks.NewDialer(proxyAddr, &shadowsocks.Config{ + Method: methodName, + PSK: pskB64, + }, nil) + + conn, err := d.DialContext(context.Background(), "tcp", "example.com:443") + if err != nil { + t.Fatalf("DialContext() error = %v", err) + } + defer conn.Close() + + if _, err := conn.Write([]byte("ping")); err != nil { + t.Fatalf("conn.Write() error = %v", err) + } + + buf := make([]byte, 4) + if _, err := io.ReadFull(conn, buf); err != nil { + t.Fatalf("conn.Read() error = %v", err) + } + if string(buf) != "pong" { + t.Fatalf("response = %q, want %q", buf, "pong") + } +} + +func TestDialer_DialContext_Deadline(t *testing.T) { + t.Parallel() + + methodName := shadowsocks.Method2022Blake3AES128GCM + pskB64 := mustBase64Key(16) + + method, err := shadowsocks.ParseMethod(methodName) + if err != nil { + t.Fatalf("ParseMethod() error = %v", err) + } + psk := decodePSKForMethod(t, method, pskB64) + + proxyAddr, stop := startMockShadowsocksServer(t, func(c net.Conn) { + defer c.Close() + + reqStart, _, err := shadowsocks.ReadTCPRequestStart(c, method, psk) + if err != nil { + return + } + + responseSalt := bytes.Repeat([]byte{0x33}, method.SaltSize) + + _, _, err = shadowsocks.WriteTCPResponseStart( + c, + method, + psk, + responseSalt, + time.Now(), + reqStart.Salt, + nil, + ) + if err != nil { + return + } + + time.Sleep(200 * time.Millisecond) + }) + defer stop() + + d := shadowsocks.NewDialer(proxyAddr, &shadowsocks.Config{ + Method: methodName, + PSK: pskB64, + }, nil) + + ctx, cancel := context.WithTimeout(context.Background(), 100*time.Millisecond) + defer cancel() + + conn, err := d.DialContext(ctx, "tcp", "example.com:443") + if err != nil { + t.Fatalf("DialContext() error = %v", err) + } + defer conn.Close() + + time.Sleep(150 * time.Millisecond) + + buf := make([]byte, 1) + _, err = conn.Read(buf) + if err == nil { + t.Fatal("expected read error after deadline") + } +} diff --git a/shadowsocks/doc.txt b/shadowsocks/doc.txt new file mode 100644 index 0000000..08e31a6 --- /dev/null +++ b/shadowsocks/doc.txt @@ -0,0 +1,307 @@ +# Shadowsocks 2022 Edition: Secure L4 Tunnel with Symmetric Encryption + +## Abstract + +This document defines the 2022 Edition of the Shadowsocks protocol. Improving upon Shadowsocks AEAD (2017), Shadowsocks 2022 addresses well-known issues of the previous editions, drops usage of obsolete cryptography, optimizes for security and performance, and leaves room for future extensions. + +## 1. Overview + +Shadowsocks 2022 is a secure proxy protocol for TCP and UDP traffic. The protocol uses [AEAD](https://en.wikipedia.org/wiki/Authenticated_encryption) with a pre-shared symmetric key to protect payload integrity and confidentiality. The proxy traffic is indistinguishable from a random byte stream, and therefore can circumvent firewalls and Internet censors that rely on [DPI (Deep Packet Inspection)](https://en.wikipedia.org/wiki/Deep_packet_inspection). + +Compared to [previous editions](https://github.com/shadowsocks/shadowsocks-org/blob/master/whitepaper/whitepaper.md) of the protocol family, Shadowsocks 2022 allows and mandates full replay protection. Each message has its unique type and cannot be used for unintended purposes. The session-based UDP proxying significantly reduces protocol overhead and improves reliability and efficiency. Obsolete cryptographic functions have been replaced by their modern counterparts. + +As with previous editions, Shadowsocks 2022 does not provide forward secrecy. It is believed that using a pre-shared key without performing handshakes is best for its use cases. + +A Shadowsocks 2022 implementation consists of a server, a client, and optionally a relay. This document specifies requirements that implementations must follow. + +### 1.1. Document Structure + +This document describes the Shadowsocks 2022 Edition and is structured as follows: + +- Section 2 describes requirements on the encryption key and how to derive session subkeys. +- Section 3 defines the encoding details of the required AES-GCM methods and the process for handling requests and responses. +- Section 4 defines the encoding details of the optional ChaCha-Poly1305 methods. + +### 1.2. Terms and Definitions + +The key words "MUST", "MUST NOT", "REQUIRED", "SHALL", "SHALL NOT", "SHOULD", "SHOULD NOT", "RECOMMENDED", "NOT RECOMMENDED", "MAY", and "OPTIONAL" in this document are to be interpreted as described in BCP 14 [RFC2119](https://www.rfc-editor.org/info/rfc2119) [RFC8174](https://www.rfc-editor.org/info/rfc8174) when, and only when, they appear in all capitals, as shown here. + +Commonly used terms in this document are described below. + +- Shadowsocks AEAD: The original AEAD construction of Shadowsocks, standardized in 2017. + +## 2. Encryption/Decryption Keys + +A pre-shared key is used to derive session subkeys, which are subsequently used to encrypt/decrypt traffic for the session. The pre-shared key is also used directly in some places. + +### 2.1. PSK + +Unlike previous editions, Shadowsocks 2022 requires that a cryptographically-secure fixed-length PSK to be directly provided by the user. Implementations MUST NOT use the old `EVP_BytesToKey` function or any other method to generate keys from passwords. + +The PSK is encoded in base64 for convenience. Practically, it can be generated with `openssl rand -base64 `. The key size depends on the chosen method. This change was inspired by WireGuard. + +| Method | Key Bytes | Salt Bytes | +| ----------------------- | --------: | ---------: | +| 2022-blake3-aes-128-gcm | 16 | 16 | +| 2022-blake3-aes-256-gcm | 32 | 32 | + +### 2.2. Subkey Derivation + +Shadowsocks 2022's subkey derivation uses [BLAKE3](https://raw.githubusercontent.com/BLAKE3-team/BLAKE3-specs/master/blake3.pdf)'s key derivation mode, which replaces the obsolete HKDF_SHA1 function in previous editions. A randomly generated salt is appended to the PSK to be used as key material. The salt has the same length as the pre-shared key. + +``` +session_subkey := blake3::derive_key(context: "shadowsocks 2022 session subkey", key_material: key + salt) +``` + +## 3. Required Methods + +Method `2022-blake3-aes-128-gcm` and `2022-blake3-aes-256-gcm` MUST be implemented by all implementations. `2022` reflects the fast-changing and flexible nature of the protocol. + +### 3.1. TCP + +TCP connections over a Shadowsocks 2022 tunnel maps 1:1 to proxy connections. Each proxy connection carries 2 proxy streams: request stream and response stream. A client initiates a proxy connection by starting a request stream, and the server sends back response over the response stream. These streams carry chunks of data encrypted by the session subkey. + +For payload transfer, Shadowsocks 2022 inherits the length-chunk-payload-chunk model from Shadowsocks AEAD, with some minor tweaks to improve performance. Standalone header chunks are added to both request and response streams to improve security and protect against replay attacks. + +#### 3.1.1. Encryption and Decryption + +Each proxy stream derives its own session subkey with a random salt for encryption and decryption. A 12-byte little-endian integer is used as nonce, and is incremented after each encryption or decryption operation. + +``` +u96le counter +aead := aead_new(key: session_subkey) +ciphertext := aead.seal(nonce: counter, plaintext) +plaintext := aead.open(nonce: counter, ciphertext) +``` + +#### 3.1.2. Format + +A request stream starts with one random salt and two standalone header chunks, followed repeatedly by one length chunk and one payload chunk. Each chunk is independently encrypted/decrypted using the AEAD cipher. + +A response stream also starts with a random salt, but only has one fixed-length header chunk, which also acts as the first length chunk. + +A length chunk is a 16-bit big-endian unsigned integer that describes the payload length in the next payload chunk. Servers and clients rely on length chunks to know how many bytes to read for the next payload chunk. + +A payload chunk can have up to 0xFFFF (65535) bytes of unencrypted payload. The 0x3FFF (16383) length cap in Shadowsocks AEAD does not apply to this edition. + +``` ++----------------+ +| length chunk | ++----------------+ +| u16 big-endian | ++----------------+ + ++---------------+ +| payload chunk | ++---------------+ +| variable | ++---------------+ + +Request stream: ++--------+------------------------+---------------------------+------------------------+---------------------------+---+ +| salt | encrypted header chunk | encrypted header chunk | encrypted length chunk | encrypted payload chunk |...| ++--------+------------------------+---------------------------+------------------------+---------------------------+---+ +| 16/32B | 11B + 16B tag | variable length + 16B tag | 2B length + 16B tag | variable length + 16B tag |...| ++--------+------------------------+---------------------------+------------------------+---------------------------+---+ + +Response stream: ++--------+------------------------+---------------------------+------------------------+---------------------------+---+ +| salt | encrypted header chunk | encrypted payload chunk | encrypted length chunk | encrypted payload chunk |...| ++--------+------------------------+---------------------------+------------------------+---------------------------+---+ +| 16/32B | 27/43B + 16B tag | variable length + 16B tag | 2B length + 16B tag | variable length + 16B tag |...| ++--------+------------------------+---------------------------+------------------------+---------------------------+---+ +``` + +#### 3.1.3. Header + +``` +Request fixed-length header: ++------+------------------+--------+ +| type | timestamp | length | ++------+------------------+--------+ +| 1B | u64be unix epoch | u16be | ++------+------------------+--------+ + +Request variable-length header: ++------+----------+-------+----------------+----------+-----------------+ +| ATYP | address | port | padding length | padding | initial payload | ++------+----------+-------+----------------+----------+-----------------+ +| 1B | variable | u16be | u16be | variable | variable | ++------+----------+-------+----------------+----------+-----------------+ + +Response fixed-length header: ++------+------------------+----------------+--------+ +| type | timestamp | request salt | length | ++------+------------------+----------------+--------+ +| 1B | u64be unix epoch | 16/32B | u16be | ++------+------------------+----------------+--------+ + +HeaderTypeClientStream = 0 +HeaderTypeServerStream = 1 +MinPaddingLength = 0 +MaxPaddingLength = 900 +``` + +- 1-byte type: Differentiates between client and server messages. A request stream has type `0`. A response stream has type `1`. +- 8-byte Unix epoch timestamp: Messages with over 30 seconds of time difference MUST be treated as replay. +- Length: Indicates the next chunk's plaintext length (not including the authentication tag). +- ATYP + address + port: Target address in [SOCKS5 address format](https://datatracker.ietf.org/doc/html/rfc1928#section-5). +- Request salt in response header: This maps a response stream to a request stream. The client MUST check this field in response header against the request salt. + +#### 3.1.3. Detection Prevention + +The random salt and header chunks MUST be buffered and sent in one write call to the underlying socket. Separate writes can result in predictable packet sizes, which could reveal the protocol in use. + +To process the salt and the fixed-length header, servers and clients MUST make exactly one read call. If the amount of data received is not enough for decryption, or decryption fails, or header validation fails, the server MUST act in a way that does not exhibit the amount of bytes consumed by the server. This defends against probes that send one byte at a time to detect how many bytes the server consumes before closing the connection. + +In such circumstances, do not immediately close the socket. Closing the socket with unread data causes RST to be sent. This reveals the exact number of bytes consumed by the server. Implementations MAY choose to employ one of the following strategies: + +1. To consistently send RST even when the receive buffer is empty, set `SO_LINGER` to true with a zero timeout, then close the socket. +2. To consistently send FIN even when the receive buffer has unread data, shut down the write half of the connection by calling `shutdown(SHUT_WR)`, then drain the connection for any further received data. +3. To consistently send FIN even when the receive buffer has unread data, but disallow unlimited writes, shut down the write half of the connection by calling `shutdown(SHUT_WR)`, then `epoll` for `EPOLLRDHUP`. Read until EOF then close the connection. This limits the amount of data the other party can send to the size of your socket receive buffer. + +In a request header, either initial payload or padding MUST be present. When making a request header, if payload is not available, add non-zero random length padding. + +For client implementations, a simple approach is to always send random length padding. To accommodate TCP Fast Open (TFO), clients MAY wait a short amount of time (typically less than one second) for client-first protocols to write the first payload, before carrying on to establish a proxy connection and write the header. + +Servers MUST reject the request if the variable-length header chunk does not contain payload and the padding length is 0. +Servers MUST enforce that the request header (including padding) does not extend beyond the header chunks. + +For response streams, the header is always sent along with payload. No padding is needed. + +#### 3.1.4. Replay Protection + +Servers MUST store all incoming salts for 60 seconds. When a new TCP session is established, the first received message is decrypted and its timestamp MUST be checked against system time. If the time difference is within 30 seconds, then the salt is checked against all stored salts. If no repeated salt is discovered, then the salt is added to the pool and the session is successfully established. + +Some techniques in implementations of previous editions are no longer necessary and SHOULD NOT be implemented for Shadowsocks 2022: + +- Clients do not need to check the salt in response streams, because the response header includes an associated request salt. +- Outgoing salts do not need to be added to the salt pool, because the header has a type field that indicates the direction of the stream. +- For salt storage, implementations MUST NOT use Bloom filters or anything that could return a false positive result, because salts only have to be stored for 60 seconds. + +### 3.2. UDP + +Shadowsocks 2022 completely overhauled UDP relay. Each UDP relay session has a unique session ID, which is also used as salt to derive the session subkey. A packet ID acts as packet counter for the session. The session ID and packet ID are combined and encrypted in a separate header. + +Clients create UDP relay sessions based on source address and port. When a client receives a packet from a new source address and port, it opens a new relay session, and subsequent packets from that source are sent over the same session. + +Servers manage UDP relay sessions by session ID. Each client session corresponds to one outgoing UDP socket on the server. + +#### 3.2.1. Encryption and Decryption + +The separate header is encrypted/decrypted with the pre-shared key using an AES block cipher. The body is encrypted/decrypted with the session subkey using an AEAD cipher specific to the session. + +``` +block_cipher := aes_new(psk) +encrypted_separate_header := block_cipher.encrypt(separate_header) +decrypted_separate_header := block_cipher.decrypt(encrypted_separate_header) + +session_subkey := blake3::derive_key(context: "shadowsocks 2022 session subkey", key_material: key + separate_header[0..8]) +session_aead_cipher := aes_gcm_new(session_subkey) +encrypted_body := session_aead_cipher.seal(nonce: separate_header[4..16], body) +decrypted_body := session_aead_cipher.open(nonce: separate_header[4..16], body) +``` + +#### 3.2.2. Format and Separate Header + +A UDP packet consists of a separate header and an AEAD-encrypted body. The separate header consists of an 8-byte session ID and an 8-byte big-endian unsigned integer as packet ID. The body is made up of the main header and payload. + +``` +Packet: ++---------------------------+---------------------------+ +| encrypted separate header | encrypted body | ++---------------------------+---------------------------+ +| 16B | variable length + 16B tag | ++---------------------------+---------------------------+ + +Separate header: ++------------+-----------+ +| session ID | packet ID | ++------------+-----------+ +| 8B | u64be | ++------------+-----------+ +``` + +UDP sessions are initiated by clients. To start a UDP session, the client generates a new random session ID and maintains a counter starting at zero as packet ID. These are used in client-to-server messages and are usually referred to as client session ID and client packet ID. + +Servers use client session IDs to identify UDP sessions. For server-to-client messages, a different set of session ID and packet ID is used, and may be referred to as server session ID and server packet ID. Like the client session ID, the server session ID MUST be randomly generated. + +#### 3.2.3. Main Header + +The main header, or message header, is the header at the start of the body. The client-to-server message header consists of type, timestamp, padding and SOCKS address. The server-to-client message header has an additional client session ID field, which maps the server session to a client session. + +``` +Client-to-server message header: ++------+------------------+----------------+----------+------+----------+-------+ +| type | timestamp | padding length | padding | ATYP | address | port | ++------+------------------+----------------+----------+------+----------+-------+ +| 1B | u64be unix epoch | u16be | variable | 1B | variable | u16be | ++------+------------------+----------------+----------+------+----------+-------+ + +Server-to-client message header: ++------+------------------+-------------------+----------------+----------+------+----------+-------+ +| type | timestamp | client session ID | padding length | padding | ATYP | address | port | ++------+------------------+-------------------+----------------+----------+------+----------+-------+ +| 1B | u64be unix epoch | 8B | u16be | variable | 1B | variable | u16be | ++------+------------------+-------------------+----------------+----------+------+----------+-------+ + +HeaderTypeClientPacket = 0 +HeaderTypeServerPacket = 1 +``` + +- 1-byte type: Differentiates between client and server messages. A client message has type `0`. A server message has type `1`. +- 8-byte Unix epoch timestamp: Messages with over 30 seconds of time difference MUST be treated as replay. +- Padding length: Specifies the length of the optional padding. Implementations MAY allow users to select from a list of predefined padding policies. Care SHOULD be taken to not exceed the network path's MTU when padding packets. + +#### 3.2.4. Session ID based Routing and Sliding Window Replay Protection + +Servers MUST route packets based on client session ID, not packet source address. When a server receives a packet with a new client session ID, a new relay session is created, and subsequent packets from that client session are sent over this relay session. + +A relay session MUST keep track of the last seen client address. When a packet is received from the client and is successfully validated, the last seen client address MUST be updated. Return packets MUST be sent to this address. This allows UDP sessions to survive client network changes. + +Each relay session MUST be remembered for at least 60 seconds. A shorter NAT timeout may allow attackers to successfully replay packets from an already forgotten client session. + +To handle server restarts, clients MUST allow each client session to be associated with more than one server session. Each association MUST be remembered for no less than the NAT timeout, which is at least 60 seconds. Alternatively, clients MAY choose to keep track of one old server session and one current server session, and reject newer server sessions when the last packet received from the old session is less than 1 minute old. + +Clients and servers MUST employ a sliding window filter for each relay session to check incoming packets for duplicate or out-of-window packet IDs. Existing implementations from WireGuard MAY be used. The packet ID MAY be checked as soon as the separate header is decrypted, but the sliding window state MUST NOT be updated before successful header validation, which filters out packets that are semantically invalid or have a bad timestamp. + +## 4. Optional Methods + +Implementations MAY choose to implement `2022-blake3-chacha20-poly1305`, `2022-blake3-chacha12-poly1305` and `2022-blake3-chacha8-poly1305` when support for CPUs without AES instructions is a priority. The use of reduced-round ChaCha20 variants is justified by [this paper](https://eprint.iacr.org/2019/1492.pdf). + +For TCP, AES-GCM is simply replaced by ChaCha-Poly1305. For UDP, a slightly different construction is used. + +### 4.1. UDP Construction + +`2022-blake3-chacha20-poly1305` uses XChaCha20-Poly1305 with the pre-shared key directly and a random nonce for each message. + +A UDP packet starts with the random nonce, followed by an encrypted body. The session ID and packet ID are merged into the main header. + +The same sliding window filter is used for replay protection. It is not necessary to check for repeated nonce. + +``` +Packet: ++-------+---------------------------+ +| nonce | encrypted body | ++-------+---------------------------+ +| 24B | variable length + 16B tag | ++-------+---------------------------+ + +Client-to-server message header: ++-------------------+------------------+------+------------------+----------------+----------+------+----------+-------+ +| client session ID | client packet ID | type | timestamp | padding length | padding | ATYP | address | port | ++-------------------+------------------+------+------------------+----------------+----------+------+----------+-------+ +| 8B | u64be | 1B | u64be unix epoch | u16be | variable | 1B | variable | u16be | ++-------------------+------------------+------+------------------+----------------+----------+------+----------+-------+ + +Server-to-client message header: ++-------------------+------------------+------+------------------+-------------------+----------------+----------+------+----------+-------+ +| server session ID | server packet ID | type | timestamp | client session ID | padding length | padding | ATYP | address | port | ++-------------------+------------------+------+------------------+-------------------+----------------+----------+------+----------+-------+ +| 8B | u64be | 1B | u64be unix epoch | 8B | u16be | variable | 1B | variable | u16be | ++-------------------+------------------+------+------------------+-------------------+----------------+----------+------+----------+-------+ +``` + +## Acknowledgement + +I would like to thank @zonyitoo, @xiaokangwang, and @nekohasekai for their input on the design of the protocol. \ No newline at end of file diff --git a/shadowsocks/key.go b/shadowsocks/key.go new file mode 100644 index 0000000..ffbe956 --- /dev/null +++ b/shadowsocks/key.go @@ -0,0 +1,78 @@ +package shadowsocks + +import ( + "crypto/rand" + "encoding/base64" + "fmt" + + "github.com/zeebo/blake3" +) + +// blake3SessionSubkeyContext is the context string used for deriving session subkeys with Blake3 in Shadowsocks 2022 methods. +const blake3SessionSubkeyContext = "shadowsocks 2022 session subkey" + +// DecodePSKTo decodes and validates a base64 PSK for the given method into dst. +// The returned slice may reuse dst's backing array. +func DecodePSKTo(dst []byte, method Method, s string) ([]byte, error) { + if err := method.Validate(); err != nil { + return nil, err + } + if s == "" { + return nil, fmt.Errorf("missing PSK") + } + + key, err := base64.StdEncoding.AppendDecode(dst[:0], []byte(s)) + if err != nil { + return nil, fmt.Errorf("invalid PSK: %w", err) + } + if len(key) != method.KeySize { + return nil, fmt.Errorf("invalid PSK length: got %d, want %d", len(key), method.KeySize) + } + + return key, nil +} + +// FillSaltTo fills dst with a random salt for the given method. +func FillSaltTo(dst []byte, method Method) error { + if err := method.Validate(); err != nil { + return err + } + if len(dst) != method.SaltSize { + return fmt.Errorf("invalid salt length: got %d, want %d", len(dst), method.SaltSize) + } + + if _, err := rand.Read(dst); err != nil { + return fmt.Errorf("generate salt: %w", err) + } + + return nil +} + +// DeriveSubkeyTo derives a session subkey into dst. +func DeriveSubkeyTo(dst []byte, method Method, key, salt []byte) error { + if err := method.Validate(); err != nil { + return err + } + if len(key) != method.KeySize { + return fmt.Errorf("invalid key length: got %d, want %d", len(key), method.KeySize) + } + if len(salt) != method.SaltSize { + return fmt.Errorf("invalid salt length: got %d, want %d", len(salt), method.SaltSize) + } + if len(dst) != method.KeySize { + return fmt.Errorf("invalid subkey length: got %d, want %d", len(dst), method.KeySize) + } + + h := blake3.NewDeriveKey(blake3SessionSubkeyContext) + if _, err := h.Write(key); err != nil { + return fmt.Errorf("derive subkey: %w", err) + } + if _, err := h.Write(salt); err != nil { + return fmt.Errorf("derive subkey: %w", err) + } + + sum := h.Sum(nil) + copy(dst, sum[:method.KeySize]) + + return nil +} diff --git a/shadowsocks/key_test.go b/shadowsocks/key_test.go new file mode 100644 index 0000000..36dd0d7 --- /dev/null +++ b/shadowsocks/key_test.go @@ -0,0 +1,259 @@ +package shadowsocks_test + +import ( + "bytes" + "encoding/base64" + "strings" + "testing" + + "github.com/33TU/socks/shadowsocks" +) + +func TestDecodePSKTo(t *testing.T) { + t.Parallel() + + method, err := shadowsocks.ParseMethod(shadowsocks.Method2022Blake3AES128GCM) + if err != nil { + t.Fatalf("ParseMethod() error = %v", err) + } + + validRaw := make([]byte, method.KeySize) + for i := range validRaw { + validRaw[i] = byte(i + 1) + } + + validPSK := base64.StdEncoding.EncodeToString(validRaw) + shortPSK := base64.StdEncoding.EncodeToString(make([]byte, method.KeySize-1)) + + tests := []struct { + name string + method shadowsocks.Method + psk string + want []byte + wantErr string + }{ + { + name: "valid", + method: method, + psk: validPSK, + want: validRaw, + }, + { + name: "missing psk", + method: method, + psk: "", + wantErr: "missing PSK", + }, + { + name: "invalid base64", + method: method, + psk: "!!!", + wantErr: "invalid PSK", + }, + { + name: "invalid key length", + method: method, + psk: shortPSK, + wantErr: "invalid PSK length", + }, + { + name: "invalid method", + method: shadowsocks.Method{ + Kind: shadowsocks.MethodKindUnknown, + }, + psk: validPSK, + wantErr: "invalid method kind", + }, + } + + for _, tt := range tests { + t.Run(tt.name, func(t *testing.T) { + t.Parallel() + + got, err := shadowsocks.DecodePSKTo(nil, tt.method, tt.psk) + if tt.wantErr != "" { + if err == nil { + t.Fatalf("expected error containing %q, got nil", tt.wantErr) + } + if !strings.Contains(err.Error(), tt.wantErr) { + t.Fatalf("expected error containing %q, got %q", tt.wantErr, err.Error()) + } + return + } + + if err != nil { + t.Fatalf("DecodePSKTo() error = %v", err) + } + if !bytes.Equal(got, tt.want) { + t.Fatalf("DecodePSKTo() = %v, want %v", got, tt.want) + } + }) + } +} + +func TestDecodePSKTo_ReuseDst(t *testing.T) { + t.Parallel() + + method, err := shadowsocks.ParseMethod(shadowsocks.Method2022Blake3AES128GCM) + if err != nil { + t.Fatalf("ParseMethod() error = %v", err) + } + + raw := make([]byte, method.KeySize) + for i := range raw { + raw[i] = byte(i + 1) + } + + psk := base64.StdEncoding.EncodeToString(raw) + dst := make([]byte, 0, 64) + + got, err := shadowsocks.DecodePSKTo(dst, method, psk) + if err != nil { + t.Fatalf("DecodePSKTo() error = %v", err) + } + + if !bytes.Equal(got, raw) { + t.Fatalf("DecodePSKTo() = %v, want %v", got, raw) + } +} + +func TestFillSalt(t *testing.T) { + t.Parallel() + + method, err := shadowsocks.ParseMethod(shadowsocks.Method2022Blake3AES256GCM) + if err != nil { + t.Fatalf("ParseMethod() error = %v", err) + } + + t.Run("valid", func(t *testing.T) { + t.Parallel() + + salt := make([]byte, method.SaltSize) + if err := shadowsocks.FillSaltTo(salt, method); err != nil { + t.Fatalf("FillSalt() error = %v", err) + } + + if len(salt) != method.SaltSize { + t.Fatalf("len(salt) = %d, want %d", len(salt), method.SaltSize) + } + + if bytes.Equal(salt, make([]byte, method.SaltSize)) { + t.Fatal("salt is all zeros, expected random data") + } + }) + + t.Run("invalid salt length", func(t *testing.T) { + t.Parallel() + + salt := make([]byte, method.SaltSize-1) + err := shadowsocks.FillSaltTo(salt, method) + if err == nil { + t.Fatal("expected error, got nil") + } + if !strings.Contains(err.Error(), "invalid salt length") { + t.Fatalf("expected invalid salt length error, got %q", err.Error()) + } + }) + + t.Run("invalid method", func(t *testing.T) { + t.Parallel() + + err := shadowsocks.FillSaltTo(make([]byte, 16), shadowsocks.Method{}) + if err == nil { + t.Fatal("expected error, got nil") + } + if !strings.Contains(err.Error(), "invalid method kind") { + t.Fatalf("expected invalid method kind error, got %q", err.Error()) + } + }) +} + +func TestDeriveSubkeyTo(t *testing.T) { + t.Parallel() + + method, err := shadowsocks.ParseMethod(shadowsocks.Method2022Blake3AES256GCM) + if err != nil { + t.Fatalf("ParseMethod() error = %v", err) + } + + key := make([]byte, method.KeySize) + salt := make([]byte, method.SaltSize) + for i := range key { + key[i] = byte(i + 1) + } + for i := range salt { + salt[i] = byte(i + 101) + } + + t.Run("valid deterministic", func(t *testing.T) { + t.Parallel() + + dst1 := make([]byte, method.KeySize) + dst2 := make([]byte, method.KeySize) + + if err := shadowsocks.DeriveSubkeyTo(dst1, method, key, salt); err != nil { + t.Fatalf("DeriveSubkeyTo() error = %v", err) + } + if err := shadowsocks.DeriveSubkeyTo(dst2, method, key, salt); err != nil { + t.Fatalf("DeriveSubkeyTo() error = %v", err) + } + + if !bytes.Equal(dst1, dst2) { + t.Fatal("derived subkeys differ for same key and salt") + } + if bytes.Equal(dst1, make([]byte, method.KeySize)) { + t.Fatal("derived subkey is all zeros") + } + }) + + t.Run("invalid key length", func(t *testing.T) { + t.Parallel() + + dst := make([]byte, method.KeySize) + err := shadowsocks.DeriveSubkeyTo(dst, method, key[:len(key)-1], salt) + if err == nil { + t.Fatal("expected error, got nil") + } + if !strings.Contains(err.Error(), "invalid key length") { + t.Fatalf("expected invalid key length error, got %q", err.Error()) + } + }) + + t.Run("invalid salt length", func(t *testing.T) { + t.Parallel() + + dst := make([]byte, method.KeySize) + err := shadowsocks.DeriveSubkeyTo(dst, method, key, salt[:len(salt)-1]) + if err == nil { + t.Fatal("expected error, got nil") + } + if !strings.Contains(err.Error(), "invalid salt length") { + t.Fatalf("expected invalid salt length error, got %q", err.Error()) + } + }) + + t.Run("invalid subkey length", func(t *testing.T) { + t.Parallel() + + dst := make([]byte, method.KeySize-1) + err := shadowsocks.DeriveSubkeyTo(dst, method, key, salt) + if err == nil { + t.Fatal("expected error, got nil") + } + if !strings.Contains(err.Error(), "invalid subkey length") { + t.Fatalf("expected invalid subkey length error, got %q", err.Error()) + } + }) + + t.Run("invalid method", func(t *testing.T) { + t.Parallel() + + err := shadowsocks.DeriveSubkeyTo(make([]byte, 16), shadowsocks.Method{}, make([]byte, 16), make([]byte, 16)) + if err == nil { + t.Fatal("expected error, got nil") + } + if !strings.Contains(err.Error(), "invalid method kind") { + t.Fatalf("expected invalid method kind error, got %q", err.Error()) + } + }) +} diff --git a/shadowsocks/method.go b/shadowsocks/method.go new file mode 100644 index 0000000..6018d9f --- /dev/null +++ b/shadowsocks/method.go @@ -0,0 +1,141 @@ +package shadowsocks + +import ( + "crypto/aes" + "crypto/cipher" + "fmt" + + "golang.org/x/crypto/chacha20poly1305" +) + +// MethodKind identifies the AEAD cipher family used by a Shadowsocks method. +type MethodKind uint8 + +const ( + MethodKindUnknown MethodKind = iota + MethodKindAESGCM + MethodKindChaCha20Poly1305 +) + +// Method describes a supported Shadowsocks method and its cryptographic parameters. +type Method struct { + Kind MethodKind + KeySize int + SaltSize int + NonceSize int + TagSize int +} + +var ( + method2022Blake3AES128GCM = Method{ + Kind: MethodKindAESGCM, + KeySize: 16, + SaltSize: 16, + NonceSize: AeadNonceSize, + TagSize: AeadTagSize, + } + + method2022Blake3AES256GCM = Method{ + Kind: MethodKindAESGCM, + KeySize: 32, + SaltSize: 32, + NonceSize: AeadNonceSize, + TagSize: AeadTagSize, + } + + method2022Blake3ChaCha20Poly1305 = Method{ + Kind: MethodKindChaCha20Poly1305, + KeySize: 32, + SaltSize: 32, + NonceSize: AeadNonceSize, + TagSize: AeadTagSize, + } +) + +// ParseMethod parses a Shadowsocks method name into its method definition. +func ParseMethod(name string) (Method, error) { + switch name { + case Method2022Blake3AES128GCM: + return method2022Blake3AES128GCM, nil + case Method2022Blake3AES256GCM: + return method2022Blake3AES256GCM, nil + case Method2022Blake3ChaCha20Poly1305: + return method2022Blake3ChaCha20Poly1305, nil + default: + return Method{}, fmt.Errorf("invalid method: %s", name) + } +} + +// IsSupportedMethod reports whether name is a supported Shadowsocks method. +func IsSupportedMethod(name string) bool { + _, err := ParseMethod(name) + return err == nil +} + +// Validate checks whether the method definition is internally valid. +func (m Method) Validate() error { + switch m.Kind { + case MethodKindAESGCM, MethodKindChaCha20Poly1305: + default: + return fmt.Errorf("invalid method kind: %d", m.Kind) + } + + if m.KeySize <= 0 { + return fmt.Errorf("invalid key size: %d", m.KeySize) + } + if m.SaltSize <= 0 { + return fmt.Errorf("invalid salt size: %d", m.SaltSize) + } + if m.NonceSize <= 0 { + return fmt.Errorf("invalid nonce size: %d", m.NonceSize) + } + if m.TagSize <= 0 { + return fmt.Errorf("invalid tag size: %d", m.TagSize) + } + + return nil +} + +// NewAEAD constructs a new AEAD instance for the method using key. +func (m Method) NewAEAD(key []byte) (cipher.AEAD, error) { + if err := m.Validate(); err != nil { + return nil, err + } + if len(key) != m.KeySize { + return nil, fmt.Errorf("invalid key length: got %d, want %d", len(key), m.KeySize) + } + + var aead cipher.AEAD + + switch m.Kind { + case MethodKindAESGCM: + block, err := aes.NewCipher(key) + if err != nil { + return nil, fmt.Errorf("create AES cipher: %w", err) + } + + aead, err = cipher.NewGCM(block) + if err != nil { + return nil, fmt.Errorf("create AES-GCM: %w", err) + } + + case MethodKindChaCha20Poly1305: + var err error + aead, err = chacha20poly1305.New(key) + if err != nil { + return nil, fmt.Errorf("create ChaCha20-Poly1305: %w", err) + } + + default: + return nil, fmt.Errorf("unsupported method kind: %d", m.Kind) + } + + if aead.NonceSize() != m.NonceSize { + return nil, fmt.Errorf("unexpected nonce size: got %d, want %d", aead.NonceSize(), m.NonceSize) + } + if aead.Overhead() != m.TagSize { + return nil, fmt.Errorf("unexpected tag size: got %d, want %d", aead.Overhead(), m.TagSize) + } + + return aead, nil +} diff --git a/shadowsocks/method_test.go b/shadowsocks/method_test.go new file mode 100644 index 0000000..9755985 --- /dev/null +++ b/shadowsocks/method_test.go @@ -0,0 +1,231 @@ +package shadowsocks_test + +import ( + "testing" + + "github.com/33TU/socks/shadowsocks" +) + +func TestParseMethod(t *testing.T) { + t.Parallel() + + tests := []struct { + name string + methodName string + wantKind shadowsocks.MethodKind + wantKeyLen int + wantErr bool + }{ + { + name: "aes128", + methodName: shadowsocks.Method2022Blake3AES128GCM, + wantKind: shadowsocks.MethodKindAESGCM, + wantKeyLen: 16, + }, + { + name: "aes256", + methodName: shadowsocks.Method2022Blake3AES256GCM, + wantKind: shadowsocks.MethodKindAESGCM, + wantKeyLen: 32, + }, + { + name: "chacha20", + methodName: shadowsocks.Method2022Blake3ChaCha20Poly1305, + wantKind: shadowsocks.MethodKindChaCha20Poly1305, + wantKeyLen: 32, + }, + { + name: "invalid", + methodName: "invalid-method", + wantErr: true, + }, + } + + for _, tt := range tests { + tt := tt + t.Run(tt.name, func(t *testing.T) { + t.Parallel() + + m, err := shadowsocks.ParseMethod(tt.methodName) + if tt.wantErr { + if err == nil { + t.Fatal("expected error, got nil") + } + return + } + if err != nil { + t.Fatalf("ParseMethod() error = %v", err) + } + + if m.Kind != tt.wantKind { + t.Fatalf("Kind = %v, want %v", m.Kind, tt.wantKind) + } + if m.KeySize != tt.wantKeyLen { + t.Fatalf("KeySize = %d, want %d", m.KeySize, tt.wantKeyLen) + } + }) + } +} + +func TestIsSupportedMethod(t *testing.T) { + t.Parallel() + + if !shadowsocks.IsSupportedMethod(shadowsocks.Method2022Blake3AES128GCM) { + t.Fatal("expected aes128 method to be supported") + } + if shadowsocks.IsSupportedMethod("invalid-method") { + t.Fatal("expected invalid method to be unsupported") + } +} + +func TestMethodValidate(t *testing.T) { + t.Parallel() + + tests := []struct { + name string + method shadowsocks.Method + wantErr bool + }{ + { + name: "valid", + method: shadowsocks.Method{ + Kind: shadowsocks.MethodKindAESGCM, + KeySize: 16, + SaltSize: 16, + NonceSize: 12, + TagSize: 16, + }, + }, + { + name: "invalid kind", + method: shadowsocks.Method{ + Kind: shadowsocks.MethodKindUnknown, + KeySize: 16, + SaltSize: 16, + NonceSize: 12, + TagSize: 16, + }, + wantErr: true, + }, + { + name: "invalid key size", + method: shadowsocks.Method{ + Kind: shadowsocks.MethodKindAESGCM, + KeySize: 0, + SaltSize: 16, + NonceSize: 12, + TagSize: 16, + }, + wantErr: true, + }, + { + name: "invalid salt size", + method: shadowsocks.Method{ + Kind: shadowsocks.MethodKindAESGCM, + KeySize: 16, + SaltSize: 0, + NonceSize: 12, + TagSize: 16, + }, + wantErr: true, + }, + { + name: "invalid nonce size", + method: shadowsocks.Method{ + Kind: shadowsocks.MethodKindAESGCM, + KeySize: 16, + SaltSize: 16, + NonceSize: 0, + TagSize: 16, + }, + wantErr: true, + }, + { + name: "invalid tag size", + method: shadowsocks.Method{ + Kind: shadowsocks.MethodKindAESGCM, + KeySize: 16, + SaltSize: 16, + NonceSize: 12, + TagSize: 0, + }, + wantErr: true, + }, + } + + for _, tt := range tests { + t.Run(tt.name, func(t *testing.T) { + t.Parallel() + + err := tt.method.Validate() + if tt.wantErr && err == nil { + t.Fatal("expected error, got nil") + } + if !tt.wantErr && err != nil { + t.Fatalf("Validate() error = %v", err) + } + }) + } +} + +func TestMethodNewAEAD(t *testing.T) { + t.Parallel() + + tests := []struct { + name string + methodName string + keyLen int + wantErr bool + }{ + { + name: "aes128", + methodName: shadowsocks.Method2022Blake3AES128GCM, + keyLen: 16, + }, + { + name: "aes256", + methodName: shadowsocks.Method2022Blake3AES256GCM, + keyLen: 32, + }, + { + name: "chacha20", + methodName: shadowsocks.Method2022Blake3ChaCha20Poly1305, + keyLen: 32, + }, + { + name: "invalid key length", + methodName: shadowsocks.Method2022Blake3AES128GCM, + keyLen: 15, + wantErr: true, + }, + } + + for _, tt := range tests { + t.Run(tt.name, func(t *testing.T) { + t.Parallel() + + m, err := shadowsocks.ParseMethod(tt.methodName) + if err != nil { + t.Fatalf("ParseMethod() error = %v", err) + } + + aead, err := m.NewAEAD(make([]byte, tt.keyLen)) + if tt.wantErr { + if err == nil { + t.Fatal("expected error, got nil") + } + return + } + if err != nil { + t.Fatalf("NewAEAD() error = %v", err) + } + + if aead.NonceSize() != m.NonceSize { + t.Fatalf("NonceSize = %d, want %d", aead.NonceSize(), m.NonceSize) + } + if aead.Overhead() != m.TagSize { + t.Fatalf("Overhead = %d, want %d", aead.Overhead(), m.TagSize) + } + }) + } +} diff --git a/shadowsocks/rand.go b/shadowsocks/rand.go new file mode 100644 index 0000000..a06a28f --- /dev/null +++ b/shadowsocks/rand.go @@ -0,0 +1,38 @@ +package shadowsocks + +import ( + "crypto/rand" + "fmt" + "math/big" +) + +// FillRandomBytes fills dst with cryptographically secure random bytes. +func FillRandomBytes(dst []byte) error { + if len(dst) == 0 { + return nil + } + n, err := rand.Read(dst) + if err != nil { + return err + } + if n != len(dst) { + return fmt.Errorf("short random read: got %d, want %d", n, len(dst)) + } + return nil +} + +// RandomInt returns a random integer in the inclusive range [min, max]. +func RandomInt(min, max int) (int, error) { + if min > max { + return 0, fmt.Errorf("invalid random range: min %d > max %d", min, max) + } + if min == max { + return min, nil + } + + n, err := rand.Int(rand.Reader, big.NewInt(int64(max-min+1))) + if err != nil { + return 0, err + } + return min + int(n.Int64()), nil +} diff --git a/shadowsocks/replay_cache.go b/shadowsocks/replay_cache.go new file mode 100644 index 0000000..3acc666 --- /dev/null +++ b/shadowsocks/replay_cache.go @@ -0,0 +1,122 @@ +package shadowsocks + +import ( + "sync" + "time" +) + +// ReplayCache stores recently seen TCP salts for replay protection. +type ReplayCache struct { + mu sync.Mutex + items map[string]time.Time +} + +// NewReplayCache creates a new replay cache. +func NewReplayCache() *ReplayCache { + return &ReplayCache{ + items: make(map[string]time.Time), + } +} + +// Seen reports whether salt is already present and not expired. +func (c *ReplayCache) Seen(salt []byte, now time.Time) bool { + if c == nil { + return false + } + + c.mu.Lock() + defer c.mu.Unlock() + + c.purgeExpiredLocked(now) + + expiry, ok := c.items[string(salt)] + if !ok { + return false + } + if !expiry.After(now) { + delete(c.items, string(salt)) + return false + } + + return true +} + +// Add stores salt until now+ttl. +func (c *ReplayCache) Add(salt []byte, now time.Time, ttl time.Duration) { + if c == nil || len(salt) == 0 || ttl <= 0 { + return + } + + c.mu.Lock() + defer c.mu.Unlock() + + c.purgeExpiredLocked(now) + c.items[string(salt)] = now.Add(ttl) +} + +// SeenOrAdd reports whether salt is already present and not expired. +// If not present, it stores the salt until now+ttl and returns false. +func (c *ReplayCache) SeenOrAdd(salt []byte, now time.Time, ttl time.Duration) bool { + if c == nil || len(salt) == 0 || ttl <= 0 { + return false + } + + c.mu.Lock() + defer c.mu.Unlock() + + c.purgeExpiredLocked(now) + + key := string(salt) + expiry, ok := c.items[key] + if ok && expiry.After(now) { + return true + } + + c.items[key] = now.Add(ttl) + return false +} + +// Remove deletes a salt from the cache. +func (c *ReplayCache) Remove(salt []byte) { + if c == nil || len(salt) == 0 { + return + } + + c.mu.Lock() + defer c.mu.Unlock() + + delete(c.items, string(salt)) +} + +// PurgeExpired removes expired entries. +func (c *ReplayCache) PurgeExpired(now time.Time) { + if c == nil { + return + } + + c.mu.Lock() + defer c.mu.Unlock() + + c.purgeExpiredLocked(now) +} + +// Len returns the number of currently stored entries after purging expired ones. +func (c *ReplayCache) Len(now time.Time) int { + if c == nil { + return 0 + } + + c.mu.Lock() + defer c.mu.Unlock() + + c.purgeExpiredLocked(now) + return len(c.items) +} + +func (c *ReplayCache) purgeExpiredLocked(now time.Time) { + for k, expiry := range c.items { + if !expiry.After(now) { + delete(c.items, k) + } + } +} diff --git a/shadowsocks/replay_cache_test.go b/shadowsocks/replay_cache_test.go new file mode 100644 index 0000000..39ed751 --- /dev/null +++ b/shadowsocks/replay_cache_test.go @@ -0,0 +1,198 @@ +package shadowsocks_test + +import ( + "testing" + "time" + + "github.com/33TU/socks/shadowsocks" +) + +func TestNewReplayCache(t *testing.T) { + t.Parallel() + + c := shadowsocks.NewReplayCache() + if c == nil { + t.Fatal("NewReplayCache() returned nil") + } + if got := c.Len(time.Now()); got != 0 { + t.Fatalf("Len() = %d, want 0", got) + } +} + +func TestReplayCache_SeenAndAdd(t *testing.T) { + t.Parallel() + + c := shadowsocks.NewReplayCache() + now := time.Unix(1700000000, 0) + salt := []byte("salt1") + + if c.Seen(salt, now) { + t.Fatal("Seen() = true before Add, want false") + } + + c.Add(salt, now, time.Minute) + + if !c.Seen(salt, now) { + t.Fatal("Seen() = false after Add, want true") + } + + if got := c.Len(now); got != 1 { + t.Fatalf("Len() = %d, want 1", got) + } +} + +func TestReplayCache_SeenOrAdd(t *testing.T) { + t.Parallel() + + c := shadowsocks.NewReplayCache() + now := time.Unix(1700000000, 0) + salt := []byte("salt1") + + if got := c.SeenOrAdd(salt, now, time.Minute); got { + t.Fatal("SeenOrAdd() first call = true, want false") + } + + if got := c.SeenOrAdd(salt, now, time.Minute); !got { + t.Fatal("SeenOrAdd() second call = false, want true") + } + + if got := c.Len(now); got != 1 { + t.Fatalf("Len() = %d, want 1", got) + } +} + +func TestReplayCache_Expiry(t *testing.T) { + t.Parallel() + + c := shadowsocks.NewReplayCache() + now := time.Unix(1700000000, 0) + salt := []byte("salt1") + + c.Add(salt, now, time.Minute) + + if !c.Seen(salt, now.Add(30*time.Second)) { + t.Fatal("Seen() before expiry = false, want true") + } + + if c.Seen(salt, now.Add(61*time.Second)) { + t.Fatal("Seen() after expiry = true, want false") + } + + if got := c.Len(now.Add(61 * time.Second)); got != 0 { + t.Fatalf("Len() after expiry = %d, want 0", got) + } +} + +func TestReplayCache_Remove(t *testing.T) { + t.Parallel() + + c := shadowsocks.NewReplayCache() + now := time.Unix(1700000000, 0) + salt := []byte("salt1") + + c.Add(salt, now, time.Minute) + c.Remove(salt) + + if c.Seen(salt, now) { + t.Fatal("Seen() after Remove = true, want false") + } + + if got := c.Len(now); got != 0 { + t.Fatalf("Len() = %d, want 0", got) + } +} + +func TestReplayCache_PurgeExpired(t *testing.T) { + t.Parallel() + + c := shadowsocks.NewReplayCache() + now := time.Unix(1700000000, 0) + + c.Add([]byte("salt1"), now, time.Minute) + c.Add([]byte("salt2"), now, 2*time.Minute) + + c.PurgeExpired(now.Add(90 * time.Second)) + + if got := c.Len(now.Add(90 * time.Second)); got != 1 { + t.Fatalf("Len() = %d, want 1", got) + } + + if c.Seen([]byte("salt1"), now.Add(90*time.Second)) { + t.Fatal("salt1 still present after expiry") + } + if !c.Seen([]byte("salt2"), now.Add(90*time.Second)) { + t.Fatal("salt2 missing before expiry") + } +} + +func TestReplayCache_NilCache(t *testing.T) { + t.Parallel() + + var c *shadowsocks.ReplayCache + now := time.Unix(1700000000, 0) + salt := []byte("salt1") + + if c.Seen(salt, now) { + t.Fatal("Seen() on nil cache = true, want false") + } + if got := c.SeenOrAdd(salt, now, time.Minute); got { + t.Fatal("SeenOrAdd() on nil cache = true, want false") + } + + c.Add(salt, now, time.Minute) + c.Remove(salt) + c.PurgeExpired(now) + + if got := c.Len(now); got != 0 { + t.Fatalf("Len() on nil cache = %d, want 0", got) + } +} + +func TestReplayCache_EmptySaltOrNonPositiveTTL(t *testing.T) { + t.Parallel() + + c := shadowsocks.NewReplayCache() + now := time.Unix(1700000000, 0) + + c.Add(nil, now, time.Minute) + c.Add([]byte{}, now, time.Minute) + c.Add([]byte("salt1"), now, 0) + c.Add([]byte("salt2"), now, -time.Second) + + if got := c.Len(now); got != 0 { + t.Fatalf("Len() = %d, want 0", got) + } + + if got := c.SeenOrAdd(nil, now, time.Minute); got { + t.Fatal("SeenOrAdd(nil) = true, want false") + } + if got := c.SeenOrAdd([]byte{}, now, time.Minute); got { + t.Fatal("SeenOrAdd(empty) = true, want false") + } + if got := c.SeenOrAdd([]byte("salt1"), now, 0); got { + t.Fatal("SeenOrAdd(ttl=0) = true, want false") + } + if got := c.SeenOrAdd([]byte("salt2"), now, -time.Second); got { + t.Fatal("SeenOrAdd(ttl<0) = true, want false") + } +} + +func TestReplayCache_SeenOrAdd_ReaddsExpired(t *testing.T) { + t.Parallel() + + c := shadowsocks.NewReplayCache() + now := time.Unix(1700000000, 0) + salt := []byte("salt1") + + if got := c.SeenOrAdd(salt, now, time.Minute); got { + t.Fatal("SeenOrAdd() first call = true, want false") + } + + if got := c.SeenOrAdd(salt, now.Add(61*time.Second), time.Minute); got { + t.Fatal("SeenOrAdd() after expiry = true, want false") + } + + if got := c.Len(now.Add(61 * time.Second)); got != 1 { + t.Fatalf("Len() = %d, want 1", got) + } +} diff --git a/shadowsocks/tcp_chunk_io.go b/shadowsocks/tcp_chunk_io.go new file mode 100644 index 0000000..2505cc5 --- /dev/null +++ b/shadowsocks/tcp_chunk_io.go @@ -0,0 +1,225 @@ +package shadowsocks + +import ( + "fmt" + "io" +) + +// TCPChunkReader reads encrypted Shadowsocks 2022 TCP chunks and exposes a +// stream-style io.Reader view over the decrypted payload. +type TCPChunkReader struct { + Cipher *TCPStreamCipher + Src io.Reader + + encLenBuf []byte + encPayloadBuf []byte + lenScratch [TcpChunkLengthLen]byte + + chunkBuf []byte + readBuf []byte +} + +// Init initializes the chunk reader for a TCP stream cipher and source. +func (r *TCPChunkReader) Init(c *TCPStreamCipher, src io.Reader) error { + if c == nil { + return fmt.Errorf("nil TCP stream cipher") + } + if src == nil { + return fmt.Errorf("nil TCP chunk source") + } + if err := c.Validate(); err != nil { + return err + } + + r.Cipher = c + r.Src = src + + encLenSize := c.EncryptedChunkLength() + if cap(r.encLenBuf) < encLenSize { + r.encLenBuf = make([]byte, encLenSize) + } else { + r.encLenBuf = r.encLenBuf[:encLenSize] + } + + r.encPayloadBuf = r.encPayloadBuf[:0] + r.chunkBuf = r.chunkBuf[:0] + r.readBuf = r.readBuf[:0] + + return nil +} + +// Validate checks whether the chunk reader is internally valid. +func (r *TCPChunkReader) Validate() error { + if r == nil { + return fmt.Errorf("nil TCP chunk reader") + } + if r.Cipher == nil { + return fmt.Errorf("missing TCP stream cipher") + } + if r.Src == nil { + return fmt.Errorf("missing TCP chunk source") + } + return r.Cipher.Validate() +} + +// ReadChunkTo reads one full encrypted TCP chunk from the underlying source, +// decrypts it, and appends the plaintext payload into dst. It returns the +// resulting slice and total bytes read from the underlying source. +func (r *TCPChunkReader) ReadChunkTo(dst []byte) ([]byte, int64, error) { + if err := r.Validate(); err != nil { + return nil, 0, err + } + + var total int64 + + n, err := io.ReadFull(r.Src, r.encLenBuf) + total += int64(n) + if err != nil { + return nil, total, err + } + + payloadLen, err := r.Cipher.DecodeChunkLength(r.encLenBuf, r.lenScratch[:0]) + if err != nil { + return nil, total, err + } + + encPayloadLen := r.Cipher.EncryptedPayloadLength(int(payloadLen)) + if cap(r.encPayloadBuf) < encPayloadLen { + r.encPayloadBuf = make([]byte, encPayloadLen) + } else { + r.encPayloadBuf = r.encPayloadBuf[:encPayloadLen] + } + + n, err = io.ReadFull(r.Src, r.encPayloadBuf) + total += int64(n) + if err != nil { + return nil, total, err + } + + dst, err = r.Cipher.DecodeChunkPayloadTo(dst, r.encPayloadBuf) + if err != nil { + return nil, total, err + } + + return dst, total, nil +} + +// Read implements io.Reader over the decrypted TCP stream. +func (r *TCPChunkReader) Read(p []byte) (int, error) { + if len(p) == 0 { + return 0, nil + } + if err := r.Validate(); err != nil { + return 0, err + } + + if len(r.readBuf) == 0 { + buf, _, err := r.ReadChunkTo(r.chunkBuf[:0]) + if err != nil { + return 0, err + } + r.chunkBuf = buf + r.readBuf = buf + } + + n := copy(p, r.readBuf) + r.readBuf = r.readBuf[n:] + return n, nil +} + +var _ io.Reader = (*TCPChunkReader)(nil) + +// TCPChunkWriter writes encrypted Shadowsocks 2022 TCP chunks and exposes a +// stream-style io.Writer interface over plaintext payload. +type TCPChunkWriter struct { + Cipher *TCPStreamCipher + Dst io.Writer + outBuf []byte +} + +// Init initializes the chunk writer for a TCP stream cipher and destination. +func (w *TCPChunkWriter) Init(c *TCPStreamCipher, dst io.Writer) error { + if c == nil { + return fmt.Errorf("nil TCP stream cipher") + } + if dst == nil { + return fmt.Errorf("nil TCP chunk destination") + } + if err := c.Validate(); err != nil { + return err + } + + w.Cipher = c + w.Dst = dst + w.outBuf = w.outBuf[:0] + return nil +} + +// Validate checks whether the chunk writer is internally valid. +func (w *TCPChunkWriter) Validate() error { + if w == nil { + return fmt.Errorf("nil TCP chunk writer") + } + if w.Cipher == nil { + return fmt.Errorf("missing TCP stream cipher") + } + if w.Dst == nil { + return fmt.Errorf("missing TCP chunk destination") + } + return w.Cipher.Validate() +} + +// WriteChunk writes one full encrypted TCP chunk to the underlying destination +// and returns the total bytes written to the underlying writer. +func (w *TCPChunkWriter) WriteChunk(payload []byte) (int64, error) { + if err := w.Validate(); err != nil { + return 0, err + } + if len(payload) > 0xFFFF { + return 0, fmt.Errorf("payload too large: got %d, max %d", len(payload), 0xFFFF) + } + + need := w.Cipher.EncryptedChunkLength() + w.Cipher.EncryptedPayloadLength(len(payload)) + if cap(w.outBuf) < need { + w.outBuf = make([]byte, 0, need) + } else { + w.outBuf = w.outBuf[:0] + } + + var err error + w.outBuf, err = w.Cipher.EncodeChunkLengthTo(w.outBuf, uint16(len(payload))) + if err != nil { + return 0, err + } + + w.outBuf, err = w.Cipher.EncodeChunkPayloadTo(w.outBuf, payload) + if err != nil { + return 0, err + } + + n, err := w.Dst.Write(w.outBuf) + return int64(n), err +} + +// Write implements io.Writer over the plaintext TCP stream. +func (w *TCPChunkWriter) Write(p []byte) (int, error) { + if len(p) == 0 { + return 0, nil + } + if err := w.Validate(); err != nil { + return 0, err + } + + written := 0 + for len(p) > 0 { + nn := min(len(p), 0xFFFF) + if _, err := w.WriteChunk(p[:nn]); err != nil { + return written, err + } + written += nn + p = p[nn:] + } + return written, nil +} + +var _ io.Writer = (*TCPChunkWriter)(nil) diff --git a/shadowsocks/tcp_chunk_io_test.go b/shadowsocks/tcp_chunk_io_test.go new file mode 100644 index 0000000..af0141d --- /dev/null +++ b/shadowsocks/tcp_chunk_io_test.go @@ -0,0 +1,418 @@ +package shadowsocks_test + +import ( + "bytes" + "io" + "strings" + "testing" + + "github.com/33TU/socks/shadowsocks" +) + +func newTCPChunkTestMethod(t *testing.T) shadowsocks.Method { + t.Helper() + + method, err := shadowsocks.ParseMethod(shadowsocks.Method2022Blake3AES128GCM) + if err != nil { + t.Fatalf("ParseMethod() error = %v", err) + } + + return method +} + +func newTCPChunkTestPSKAndSalt(method shadowsocks.Method) ([]byte, []byte) { + psk := make([]byte, method.KeySize) + salt := make([]byte, method.SaltSize) + + for i := range psk { + psk[i] = byte(i + 1) + } + for i := range salt { + salt[i] = byte(i + 101) + } + + return psk, salt +} + +func newTCPChunkCipherPair(t *testing.T) (*shadowsocks.TCPStreamCipher, *shadowsocks.TCPStreamCipher) { + t.Helper() + + method := newTCPChunkTestMethod(t) + psk, salt := newTCPChunkTestPSKAndSalt(method) + + enc, err := shadowsocks.NewTCPStreamCipherFromPSK(method, psk, salt) + if err != nil { + t.Fatalf("NewTCPStreamCipherFromPSK(enc) error = %v", err) + } + + dec, err := shadowsocks.NewTCPStreamCipherFromPSK(method, psk, salt) + if err != nil { + t.Fatalf("NewTCPStreamCipherFromPSK(dec) error = %v", err) + } + + return enc, dec +} + +func TestTCPChunkReader_Init_Validate(t *testing.T) { + t.Parallel() + + _, dec := newTCPChunkCipherPair(t) + + t.Run("valid", func(t *testing.T) { + t.Parallel() + + var r shadowsocks.TCPChunkReader + if err := r.Init(dec, bytes.NewReader(nil)); err != nil { + t.Fatalf("Init() error = %v", err) + } + if err := r.Validate(); err != nil { + t.Fatalf("Validate() error = %v", err) + } + }) + + t.Run("nil cipher", func(t *testing.T) { + t.Parallel() + + var r shadowsocks.TCPChunkReader + err := r.Init(nil, bytes.NewReader(nil)) + if err == nil { + t.Fatal("expected error, got nil") + } + if !strings.Contains(err.Error(), "nil TCP stream cipher") { + t.Fatalf("unexpected error: %v", err) + } + }) + + t.Run("nil source", func(t *testing.T) { + t.Parallel() + + var r shadowsocks.TCPChunkReader + err := r.Init(dec, nil) + if err == nil { + t.Fatal("expected error, got nil") + } + if !strings.Contains(err.Error(), "nil TCP chunk source") { + t.Fatalf("unexpected error: %v", err) + } + }) + + t.Run("nil reader validate", func(t *testing.T) { + t.Parallel() + + var r *shadowsocks.TCPChunkReader + err := r.Validate() + if err == nil { + t.Fatal("expected error, got nil") + } + if !strings.Contains(err.Error(), "nil TCP chunk reader") { + t.Fatalf("unexpected error: %v", err) + } + }) + + t.Run("missing cipher", func(t *testing.T) { + t.Parallel() + + r := &shadowsocks.TCPChunkReader{Src: bytes.NewReader(nil)} + err := r.Validate() + if err == nil { + t.Fatal("expected error, got nil") + } + if !strings.Contains(err.Error(), "missing TCP stream cipher") { + t.Fatalf("unexpected error: %v", err) + } + }) + + t.Run("missing source", func(t *testing.T) { + t.Parallel() + + r := &shadowsocks.TCPChunkReader{Cipher: dec} + err := r.Validate() + if err == nil { + t.Fatal("expected error, got nil") + } + if !strings.Contains(err.Error(), "missing TCP chunk source") { + t.Fatalf("unexpected error: %v", err) + } + }) +} + +func TestTCPChunkWriter_Init_Validate(t *testing.T) { + t.Parallel() + + enc, _ := newTCPChunkCipherPair(t) + + t.Run("valid", func(t *testing.T) { + t.Parallel() + + var w shadowsocks.TCPChunkWriter + if err := w.Init(enc, io.Discard); err != nil { + t.Fatalf("Init() error = %v", err) + } + if err := w.Validate(); err != nil { + t.Fatalf("Validate() error = %v", err) + } + }) + + t.Run("nil cipher", func(t *testing.T) { + t.Parallel() + + var w shadowsocks.TCPChunkWriter + err := w.Init(nil, io.Discard) + if err == nil { + t.Fatal("expected error, got nil") + } + if !strings.Contains(err.Error(), "nil TCP stream cipher") { + t.Fatalf("unexpected error: %v", err) + } + }) + + t.Run("nil destination", func(t *testing.T) { + t.Parallel() + + var w shadowsocks.TCPChunkWriter + err := w.Init(enc, nil) + if err == nil { + t.Fatal("expected error, got nil") + } + if !strings.Contains(err.Error(), "nil TCP chunk destination") { + t.Fatalf("unexpected error: %v", err) + } + }) + + t.Run("nil writer validate", func(t *testing.T) { + t.Parallel() + + var w *shadowsocks.TCPChunkWriter + err := w.Validate() + if err == nil { + t.Fatal("expected error, got nil") + } + if !strings.Contains(err.Error(), "nil TCP chunk writer") { + t.Fatalf("unexpected error: %v", err) + } + }) + + t.Run("missing cipher", func(t *testing.T) { + t.Parallel() + + w := &shadowsocks.TCPChunkWriter{Dst: io.Discard} + err := w.Validate() + if err == nil { + t.Fatal("expected error, got nil") + } + if !strings.Contains(err.Error(), "missing TCP stream cipher") { + t.Fatalf("unexpected error: %v", err) + } + }) + + t.Run("missing destination", func(t *testing.T) { + t.Parallel() + + w := &shadowsocks.TCPChunkWriter{Cipher: enc} + err := w.Validate() + if err == nil { + t.Fatal("expected error, got nil") + } + if !strings.Contains(err.Error(), "missing TCP chunk destination") { + t.Fatalf("unexpected error: %v", err) + } + }) +} + +func TestTCPChunkWriterReader_RoundTrip(t *testing.T) { + t.Parallel() + + enc, dec := newTCPChunkCipherPair(t) + + var buf bytes.Buffer + + var w shadowsocks.TCPChunkWriter + if err := w.Init(enc, &buf); err != nil { + t.Fatalf("writer Init() error = %v", err) + } + + var r shadowsocks.TCPChunkReader + if err := r.Init(dec, &buf); err != nil { + t.Fatalf("reader Init() error = %v", err) + } + + payload := []byte("hello chunk") + + nw, err := w.WriteChunk(payload) + if err != nil { + t.Fatalf("WriteChunk() error = %v", err) + } + if nw != int64(buf.Len()) { + t.Fatalf("WriteChunk() wrote %d bytes, buffer has %d", nw, buf.Len()) + } + + got, nr, err := r.ReadChunkTo(nil) + if err != nil { + t.Fatalf("ReadChunkTo() error = %v", err) + } + if nr != nw { + t.Fatalf("ReadChunkTo() read %d bytes, want %d", nr, nw) + } + if !bytes.Equal(got, payload) { + t.Fatalf("ReadChunkTo() = %q, want %q", got, payload) + } +} + +func TestTCPChunkReader_ReadChunkTo_AppendsToDst(t *testing.T) { + t.Parallel() + + enc, dec := newTCPChunkCipherPair(t) + + var buf bytes.Buffer + + var w shadowsocks.TCPChunkWriter + if err := w.Init(enc, &buf); err != nil { + t.Fatalf("writer Init() error = %v", err) + } + + var r shadowsocks.TCPChunkReader + if err := r.Init(dec, &buf); err != nil { + t.Fatalf("reader Init() error = %v", err) + } + + payload := []byte("payload") + prefix := []byte("prefix:") + + if _, err := w.WriteChunk(payload); err != nil { + t.Fatalf("WriteChunk() error = %v", err) + } + + got, _, err := r.ReadChunkTo(append([]byte(nil), prefix...)) + if err != nil { + t.Fatalf("ReadChunkTo() error = %v", err) + } + + want := append(append([]byte(nil), prefix...), payload...) + if !bytes.Equal(got, want) { + t.Fatalf("ReadChunkTo() = %q, want %q", got, want) + } +} + +func TestTCPChunkWriter_WriteChunk_TooLarge(t *testing.T) { + t.Parallel() + + enc, _ := newTCPChunkCipherPair(t) + + var w shadowsocks.TCPChunkWriter + if err := w.Init(enc, io.Discard); err != nil { + t.Fatalf("Init() error = %v", err) + } + + payload := make([]byte, 0x10000) + _, err := w.WriteChunk(payload) + if err == nil { + t.Fatal("expected error, got nil") + } + if !strings.Contains(err.Error(), "payload too large") { + t.Fatalf("unexpected error: %v", err) + } +} + +func TestTCPChunkReader_ReadChunk_ShortLengthRead(t *testing.T) { + t.Parallel() + + _, dec := newTCPChunkCipherPair(t) + + short := bytes.NewReader([]byte{1, 2, 3}) + + var r shadowsocks.TCPChunkReader + if err := r.Init(dec, short); err != nil { + t.Fatalf("Init() error = %v", err) + } + + _, n, err := r.ReadChunkTo(nil) + if err == nil { + t.Fatal("expected error, got nil") + } + if n != 3 { + t.Fatalf("ReadChunkTo() read %d bytes, want 3", n) + } + if err != io.EOF && err != io.ErrUnexpectedEOF { + t.Fatalf("unexpected error: %v", err) + } +} + +func TestTCPChunkReader_ReadChunk_ShortPayloadRead(t *testing.T) { + t.Parallel() + + enc, dec := newTCPChunkCipherPair(t) + + var buf bytes.Buffer + + var w shadowsocks.TCPChunkWriter + if err := w.Init(enc, &buf); err != nil { + t.Fatalf("writer Init() error = %v", err) + } + + payload := []byte("hello payload") + nw, err := w.WriteChunk(payload) + if err != nil { + t.Fatalf("WriteChunk() error = %v", err) + } + + wire := buf.Bytes() + shortWire := wire[:len(wire)-2] + + var r shadowsocks.TCPChunkReader + if err := r.Init(dec, bytes.NewReader(shortWire)); err != nil { + t.Fatalf("reader Init() error = %v", err) + } + + _, nr, err := r.ReadChunkTo(nil) + if err == nil { + t.Fatal("expected error, got nil") + } + if nr != nw-2 { + t.Fatalf("ReadChunkTo() read %d bytes, want %d", nr, nw-2) + } + if err != io.EOF && err != io.ErrUnexpectedEOF { + t.Fatalf("unexpected error: %v", err) + } +} + +func TestTCPChunkReader_ReadChunk_WrongCipherFails(t *testing.T) { + t.Parallel() + + enc, _ := newTCPChunkCipherPair(t) + + method := newTCPChunkTestMethod(t) + psk := make([]byte, method.KeySize) + salt := make([]byte, method.SaltSize) + for i := range psk { + psk[i] = byte(i + 9) + } + for i := range salt { + salt[i] = byte(i + 19) + } + + wrongDec, err := shadowsocks.NewTCPStreamCipherFromPSK(method, psk, salt) + if err != nil { + t.Fatalf("NewTCPStreamCipherFromPSK() error = %v", err) + } + + var buf bytes.Buffer + + var w shadowsocks.TCPChunkWriter + if err := w.Init(enc, &buf); err != nil { + t.Fatalf("writer Init() error = %v", err) + } + + var r shadowsocks.TCPChunkReader + if err := r.Init(wrongDec, &buf); err != nil { + t.Fatalf("reader Init() error = %v", err) + } + + if _, err := w.WriteChunk([]byte("hello")); err != nil { + t.Fatalf("WriteChunk() error = %v", err) + } + + _, _, err = r.ReadChunkTo(nil) + if err == nil { + t.Fatal("expected decrypt error, got nil") + } +} diff --git a/shadowsocks/tcp_conn.go b/shadowsocks/tcp_conn.go new file mode 100644 index 0000000..c2da6da --- /dev/null +++ b/shadowsocks/tcp_conn.go @@ -0,0 +1,208 @@ +package shadowsocks + +import ( + "fmt" + "net" + "time" +) + +type TCPConn struct { + net.Conn + + Reader TCPChunkReader + Writer TCPChunkWriter + + responseMethod Method + responsePSK []byte + requestSalt []byte +} + +func NewClientTCPConn( + conn net.Conn, + method Method, + psk []byte, + target Addr, + padding []byte, + initialPayload []byte, +) (*TCPConn, error) { + if conn == nil { + return nil, fmt.Errorf("nil net.Conn") + } + if err := method.Validate(); err != nil { + return nil, err + } + if len(psk) != method.KeySize { + return nil, fmt.Errorf("invalid PSK length: got %d, want %d", len(psk), method.KeySize) + } + if err := target.Validate(); err != nil { + return nil, err + } + + requestSalt := make([]byte, method.SaltSize) + if err := FillSaltTo(requestSalt, method); err != nil { + return nil, err + } + + requestCipher, _, err := WriteTCPRequestStart( + conn, + method, + psk, + requestSalt, + time.Now(), + target, + padding, + initialPayload, + ) + if err != nil { + return nil, err + } + + var writer TCPChunkWriter + if err := writer.Init(requestCipher, conn); err != nil { + return nil, err + } + + return &TCPConn{ + Conn: conn, + Writer: writer, + responseMethod: method, + responsePSK: append([]byte(nil), psk...), + requestSalt: append([]byte(nil), requestSalt...), + }, nil +} + +func NewServerTCPConn(conn net.Conn, method Method, psk []byte) (*TCPConn, *ParsedTCPRequestStart, error) { + if conn == nil { + return nil, nil, fmt.Errorf("nil net.Conn") + } + if err := method.Validate(); err != nil { + return nil, nil, err + } + if len(psk) != method.KeySize { + return nil, nil, fmt.Errorf("invalid PSK length: got %d, want %d", len(psk), method.KeySize) + } + + reqStart, _, err := ReadTCPRequestStart(conn, method, psk) + if err != nil { + return nil, nil, err + } + + var reader TCPChunkReader + if err := reader.Init(reqStart.Cipher, conn); err != nil { + return nil, nil, err + } + + c := &TCPConn{ + Conn: conn, + Reader: reader, + } + + if len(reqStart.Header.InitialData) > 0 { + c.Reader.chunkBuf = append(c.Reader.chunkBuf[:0], reqStart.Header.InitialData...) + c.Reader.readBuf = c.Reader.chunkBuf + } + + return c, reqStart, nil +} + +func (c *TCPConn) InitResponse(method Method, psk []byte, requestSalt []byte, initialPayload []byte) error { + if c == nil { + return fmt.Errorf("nil TCPConn") + } + if c.Conn == nil { + return fmt.Errorf("nil net.Conn") + } + if err := method.Validate(); err != nil { + return err + } + if len(psk) != method.KeySize { + return fmt.Errorf("invalid PSK length: got %d, want %d", len(psk), method.KeySize) + } + if len(requestSalt) != method.SaltSize { + return fmt.Errorf("invalid request salt length: got %d, want %d", len(requestSalt), method.SaltSize) + } + + responseSalt := make([]byte, method.SaltSize) + if err := FillSaltTo(responseSalt, method); err != nil { + return err + } + + responseCipher, _, err := WriteTCPResponseStart( + c.Conn, + method, + psk, + responseSalt, + time.Now(), + requestSalt, + initialPayload, + ) + if err != nil { + return err + } + + return c.Writer.Init(responseCipher, c.Conn) +} + +func (c *TCPConn) ensureClientResponseReady() error { + if c == nil { + return fmt.Errorf("nil TCPConn") + } + if c.Conn == nil { + return fmt.Errorf("nil net.Conn") + } + if c.Reader.Cipher != nil { + return nil + } + if err := c.responseMethod.Validate(); err != nil { + return err + } + if len(c.responsePSK) != c.responseMethod.KeySize { + return fmt.Errorf("invalid response PSK length: got %d, want %d", len(c.responsePSK), c.responseMethod.KeySize) + } + if len(c.requestSalt) != c.responseMethod.SaltSize { + return fmt.Errorf("invalid request salt length: got %d, want %d", len(c.requestSalt), c.responseMethod.SaltSize) + } + + respStart, _, err := ReadTCPResponseStart(c.Conn, c.responseMethod, c.responsePSK, c.requestSalt) + if err != nil { + return err + } + + if err := c.Reader.Init(respStart.Cipher, c.Conn); err != nil { + return err + } + + if len(respStart.InitialPayload) > 0 { + c.Reader.chunkBuf = append(c.Reader.chunkBuf[:0], respStart.InitialPayload...) + c.Reader.readBuf = c.Reader.chunkBuf + } + + return nil +} + +func (c *TCPConn) Read(p []byte) (int, error) { + if len(p) == 0 { + return 0, nil + } + if c == nil { + return 0, fmt.Errorf("nil TCPConn") + } + if c.Reader.Cipher == nil { + if err := c.ensureClientResponseReady(); err != nil { + return 0, err + } + } + return c.Reader.Read(p) +} + +func (c *TCPConn) Write(p []byte) (int, error) { + if len(p) == 0 { + return 0, nil + } + if c == nil { + return 0, fmt.Errorf("nil TCPConn") + } + return c.Writer.Write(p) +} + +var _ net.Conn = (*TCPConn)(nil) diff --git a/shadowsocks/tcp_conn_test.go b/shadowsocks/tcp_conn_test.go new file mode 100644 index 0000000..ab64749 --- /dev/null +++ b/shadowsocks/tcp_conn_test.go @@ -0,0 +1,335 @@ +package shadowsocks_test + +import ( + "bytes" + "errors" + "net" + "testing" + + "github.com/33TU/socks/shadowsocks" +) + +func newTCPConnTestMethod(t *testing.T) shadowsocks.Method { + t.Helper() + + method, err := shadowsocks.ParseMethod(shadowsocks.Method2022Blake3AES128GCM) + if err != nil { + t.Fatalf("ParseMethod() error = %v", err) + } + + return method +} + +func newTCPConnTestPSKAndSalt(method shadowsocks.Method) ([]byte, []byte) { + psk := make([]byte, method.KeySize) + salt := make([]byte, method.SaltSize) + + for i := range psk { + psk[i] = byte(i + 1) + } + for i := range salt { + salt[i] = byte(i + 101) + } + + return psk, salt +} + +func newTCPConnCipherPair(t *testing.T) (*shadowsocks.TCPStreamCipher, *shadowsocks.TCPStreamCipher) { + t.Helper() + + method := newTCPConnTestMethod(t) + psk, salt := newTCPConnTestPSKAndSalt(method) + + enc, err := shadowsocks.NewTCPStreamCipherFromPSK(method, psk, salt) + if err != nil { + t.Fatalf("NewTCPStreamCipherFromPSK(enc) error = %v", err) + } + + dec, err := shadowsocks.NewTCPStreamCipherFromPSK(method, psk, salt) + if err != nil { + t.Fatalf("NewTCPStreamCipherFromPSK(dec) error = %v", err) + } + + return enc, dec +} + +func TestTCPConn_Read(t *testing.T) { + t.Parallel() + + serverConn, clientConn := net.Pipe() + defer serverConn.Close() + defer clientConn.Close() + + enc, dec := newTCPConnCipherPair(t) + + var writer shadowsocks.TCPChunkWriter + if err := writer.Init(enc, serverConn); err != nil { + t.Fatalf("writer.Init() error = %v", err) + } + + var reader shadowsocks.TCPChunkReader + if err := reader.Init(dec, clientConn); err != nil { + t.Fatalf("reader.Init() error = %v", err) + } + + c := &shadowsocks.TCPConn{ + Conn: clientConn, + Reader: reader, + Writer: readerlessWriter(enc, clientConn), + } + + want := []byte("hello world") + + go func() { + defer serverConn.Close() + if _, err := writer.WriteChunk(want); err != nil { + t.Errorf("WriteChunk() error = %v", err) + } + }() + + buf := make([]byte, len(want)) + n, err := c.Read(buf) + if err != nil { + t.Fatalf("Read() error = %v", err) + } + if n != len(want) { + t.Fatalf("Read() n = %d, want %d", n, len(want)) + } + if !bytes.Equal(buf[:n], want) { + t.Fatalf("Read() = %q, want %q", buf[:n], want) + } +} + +func TestTCPConn_Read_PartialBuffered(t *testing.T) { + t.Parallel() + + serverConn, clientConn := net.Pipe() + defer serverConn.Close() + defer clientConn.Close() + + enc, dec := newTCPConnCipherPair(t) + + var writer shadowsocks.TCPChunkWriter + if err := writer.Init(enc, serverConn); err != nil { + t.Fatalf("writer.Init() error = %v", err) + } + + var reader shadowsocks.TCPChunkReader + if err := reader.Init(dec, clientConn); err != nil { + t.Fatalf("reader.Init() error = %v", err) + } + + c := &shadowsocks.TCPConn{ + Conn: clientConn, + Reader: reader, + Writer: readerlessWriter(enc, clientConn), + } + + want := []byte("hello world") + + go func() { + defer serverConn.Close() + if _, err := writer.WriteChunk(want); err != nil { + t.Errorf("WriteChunk() error = %v", err) + } + }() + + buf1 := make([]byte, 5) + n1, err := c.Read(buf1) + if err != nil { + t.Fatalf("first Read() error = %v", err) + } + if n1 != 5 { + t.Fatalf("first Read() n = %d, want 5", n1) + } + if !bytes.Equal(buf1[:n1], []byte("hello")) { + t.Fatalf("first Read() = %q, want %q", buf1[:n1], "hello") + } + + buf2 := make([]byte, 6) + n2, err := c.Read(buf2) + if err != nil { + t.Fatalf("second Read() error = %v", err) + } + if n2 != 6 { + t.Fatalf("second Read() n = %d, want 6", n2) + } + if !bytes.Equal(buf2[:n2], []byte(" world")) { + t.Fatalf("second Read() = %q, want %q", buf2[:n2], " world") + } +} + +func TestTCPConn_Write(t *testing.T) { + t.Parallel() + + serverConn, clientConn := net.Pipe() + defer serverConn.Close() + defer clientConn.Close() + + enc, dec := newTCPConnCipherPair(t) + + var writer shadowsocks.TCPChunkWriter + if err := writer.Init(enc, clientConn); err != nil { + t.Fatalf("writer.Init() error = %v", err) + } + + var reader shadowsocks.TCPChunkReader + if err := reader.Init(dec, serverConn); err != nil { + t.Fatalf("reader.Init() error = %v", err) + } + + c := &shadowsocks.TCPConn{ + Conn: clientConn, + Reader: writerlessReader(dec, clientConn), + Writer: writer, + } + + want := []byte("hello world") + errCh := make(chan error, 1) + + go func() { + defer serverConn.Close() + + got, _, err := reader.ReadChunkTo(nil) + if err != nil { + errCh <- err + return + } + if !bytes.Equal(got, want) { + errCh <- errors.New("payload mismatch") + return + } + errCh <- nil + }() + + n, err := c.Write(want) + if err != nil { + t.Fatalf("Write() error = %v", err) + } + if n != len(want) { + t.Fatalf("Write() n = %d, want %d", n, len(want)) + } + + if err := <-errCh; err != nil { + t.Fatalf("server read error = %v", err) + } +} + +func TestTCPConn_Write_Empty(t *testing.T) { + t.Parallel() + + serverConn, clientConn := net.Pipe() + defer serverConn.Close() + defer clientConn.Close() + + enc, dec := newTCPConnCipherPair(t) + + var writer shadowsocks.TCPChunkWriter + if err := writer.Init(enc, clientConn); err != nil { + t.Fatalf("writer.Init() error = %v", err) + } + + var reader shadowsocks.TCPChunkReader + if err := reader.Init(dec, clientConn); err != nil { + t.Fatalf("reader.Init() error = %v", err) + } + + c := &shadowsocks.TCPConn{ + Conn: clientConn, + Reader: reader, + Writer: writer, + } + + n, err := c.Write(nil) + if err != nil { + t.Fatalf("Write() error = %v", err) + } + if n != 0 { + t.Fatalf("Write() n = %d, want 0", n) + } +} + +func TestTCPConn_Write_SplitsLargePayload(t *testing.T) { + t.Parallel() + + serverConn, clientConn := net.Pipe() + defer serverConn.Close() + defer clientConn.Close() + + enc, dec := newTCPConnCipherPair(t) + + var writer shadowsocks.TCPChunkWriter + if err := writer.Init(enc, clientConn); err != nil { + t.Fatalf("writer.Init() error = %v", err) + } + + var reader shadowsocks.TCPChunkReader + if err := reader.Init(dec, serverConn); err != nil { + t.Fatalf("reader.Init() error = %v", err) + } + + c := &shadowsocks.TCPConn{ + Conn: clientConn, + Reader: writerlessReader(dec, clientConn), + Writer: writer, + } + + want := bytes.Repeat([]byte("a"), 0x10000+123) + done := make(chan error, 1) + + go func() { + defer serverConn.Close() + + var got []byte + + part1, _, err := reader.ReadChunkTo(nil) + if err != nil { + done <- err + return + } + got = append(got, part1...) + + part2, _, err := reader.ReadChunkTo(nil) + if err != nil { + done <- err + return + } + got = append(got, part2...) + + if !bytes.Equal(got, want) { + done <- errors.New("payload mismatch") + return + } + done <- nil + }() + + n, err := c.Write(want) + if err != nil { + t.Fatalf("Write() error = %v", err) + } + if n != len(want) { + t.Fatalf("Write() n = %d, want %d", n, len(want)) + } + + if err := <-done; err != nil { + t.Fatalf("server read error = %v", err) + } +} + +// helpers for constructing a TCPConn in tests where only one side is actually used + +func readerlessWriter(enc *shadowsocks.TCPStreamCipher, dst net.Conn) shadowsocks.TCPChunkWriter { + var w shadowsocks.TCPChunkWriter + if err := w.Init(enc, dst); err != nil { + panic(err) + } + return w +} + +func writerlessReader(dec *shadowsocks.TCPStreamCipher, src net.Conn) shadowsocks.TCPChunkReader { + var r shadowsocks.TCPChunkReader + if err := r.Init(dec, src); err != nil { + panic(err) + } + return r +} diff --git a/shadowsocks/tcp_proto.go b/shadowsocks/tcp_proto.go new file mode 100644 index 0000000..2ec7153 --- /dev/null +++ b/shadowsocks/tcp_proto.go @@ -0,0 +1,33 @@ +package shadowsocks + +import ( + "errors" +) + +// TCP header types for Shadowsocks 2022 stream protocol. +const ( + TCPHeaderTypeClientStream = 0x00 + TCPHeaderTypeServerStream = 0x01 +) + +const ( + TcpRequestFixedHeaderLen = 1 + 8 + 2 + TcpResponseFixedBaseLen = 1 + 8 + 2 +) + +const ( + AeadNonceSize = 12 + AeadTagSize = 16 +) + +const TcpChunkLengthLen = 2 + +// Common validation and decode errors for Shadowsocks TCP headers. +var ( + ErrInvalidTCPHeaderType = errors.New("invalid TCP header type") // res and req + ErrInvalidTCPPaddingLength = errors.New("invalid TCP padding length") // req + ErrMissingTCPHeaderData = errors.New("missing TCP header data") // req + ErrShortTCPHeader = errors.New("short TCP header") // res and req + ErrMissingTCPResponseSalt = errors.New("missing TCP response salt") // res + ErrInvalidTCPResponseSaltLen = errors.New("invalid TCP response salt length") // res +) diff --git a/shadowsocks/tcp_request_header.go b/shadowsocks/tcp_request_header.go new file mode 100644 index 0000000..4d5a87e --- /dev/null +++ b/shadowsocks/tcp_request_header.go @@ -0,0 +1,176 @@ +package shadowsocks + +import ( + "encoding/binary" + "fmt" +) + +// TCPRequestFixedHeader represents the fixed-length request header used by +// Shadowsocks 2022 TCP streams. +type TCPRequestFixedHeader struct { + Type byte + Timestamp uint64 + Length uint16 +} + +// Init initializes a TCPRequestFixedHeader. +func (h *TCPRequestFixedHeader) Init(typ byte, timestamp uint64, length uint16) { + h.Type = typ + h.Timestamp = timestamp + h.Length = length +} + +// Validate checks the correctness of the fixed request header fields. +func (h *TCPRequestFixedHeader) Validate() error { + if h.Type != TCPHeaderTypeClientStream { + return ErrInvalidTCPHeaderType + } + + return nil +} + +// EncodedLen returns the number of bytes required to encode the fixed request header. +func (h *TCPRequestFixedHeader) EncodedLen() int { + return TcpRequestFixedHeaderLen +} + +// Decode decodes a fixed request header from src. +// It returns the number of bytes consumed. +func (h *TCPRequestFixedHeader) Decode(src []byte) (int, error) { + if len(src) < TcpRequestFixedHeaderLen { + return 0, ErrShortTCPHeader + } + + h.Type = src[0] + h.Timestamp = binary.BigEndian.Uint64(src[1:9]) + h.Length = binary.BigEndian.Uint16(src[9:11]) + + return TcpRequestFixedHeaderLen, h.Validate() +} + +// EncodeTo encodes the fixed request header into dst. +func (h *TCPRequestFixedHeader) EncodeTo(dst []byte) ([]byte, error) { + if err := h.Validate(); err != nil { + return nil, err + } + + dst = append(dst, h.Type) + dst = binary.BigEndian.AppendUint64(dst, h.Timestamp) + dst = binary.BigEndian.AppendUint16(dst, h.Length) + + return dst, nil +} + +// String returns a human-readable representation of the fixed request header. +func (h *TCPRequestFixedHeader) String() string { + return fmt.Sprintf( + "TCPRequestFixedHeader{Type:%d Timestamp:%d Length:%d}", + h.Type, h.Timestamp, h.Length, + ) +} + +//////// + +// TCPRequestVariableHeader represents the variable-length request header used by +// Shadowsocks 2022 TCP streams. +type TCPRequestVariableHeader struct { + Target Addr + PaddingLen uint16 + Padding []byte + InitialData []byte +} + +// Init initializes a TCPRequestVariableHeader. +func (h *TCPRequestVariableHeader) Init(target Addr, padding, initialData []byte) { + h.Target = target + h.PaddingLen = uint16(len(padding)) + h.Padding = padding + h.InitialData = initialData +} + +// Validate checks the correctness of the variable request header fields. +func (h *TCPRequestVariableHeader) Validate() error { + if err := h.Target.Validate(); err != nil { + return err + } + if int(h.PaddingLen) != len(h.Padding) { + return ErrInvalidTCPPaddingLength + } + if len(h.Padding) == 0 && len(h.InitialData) == 0 { + return ErrMissingTCPHeaderData + } + + return nil +} + +// EncodedLen returns the number of bytes required to encode the variable request header. +func (h *TCPRequestVariableHeader) EncodedLen() int { + return h.Target.EncodedLen() + 2 + len(h.Padding) + len(h.InitialData) +} + +// Decode decodes a variable request header from src. +// It returns the number of bytes consumed. +func (h *TCPRequestVariableHeader) Decode(src []byte) (int, error) { + h.Target = Addr{} + h.PaddingLen = 0 + h.Padding = nil + h.InitialData = nil + + n, err := h.Target.Decode(src) + if err != nil { + return 0, err + } + if len(src[n:]) < 2 { + return 0, ErrShortTCPHeader + } + + h.PaddingLen = binary.BigEndian.Uint16(src[n : n+2]) + n += 2 + + if len(src[n:]) < int(h.PaddingLen) { + return 0, ErrShortTCPHeader + } + + if h.PaddingLen > 0 { + h.Padding = append(h.Padding[:0], src[n:n+int(h.PaddingLen)]...) + n += int(h.PaddingLen) + } + + if len(src[n:]) > 0 { + h.InitialData = append(h.InitialData[:0], src[n:]...) + n = len(src) + } + + if err := h.Validate(); err != nil { + return 0, err + } + + return n, nil +} + +// EncodeTo encodes the variable request header into dst. +func (h *TCPRequestVariableHeader) EncodeTo(dst []byte) ([]byte, error) { + if err := h.Validate(); err != nil { + return nil, err + } + + var err error + dst, err = h.Target.EncodeTo(dst) + if err != nil { + return nil, err + } + + dst = binary.BigEndian.AppendUint16(dst, h.PaddingLen) + dst = append(dst, h.Padding...) + dst = append(dst, h.InitialData...) + + return dst, nil +} + +// String returns a human-readable representation of the variable request header. +func (h *TCPRequestVariableHeader) String() string { + return fmt.Sprintf( + "TCPRequestVariableHeader{Target:%s PaddingLen:%d InitialDataLen:%d}", + h.Target.String(), h.PaddingLen, len(h.InitialData), + ) +} diff --git a/shadowsocks/tcp_request_header_test.go b/shadowsocks/tcp_request_header_test.go new file mode 100644 index 0000000..1120d17 --- /dev/null +++ b/shadowsocks/tcp_request_header_test.go @@ -0,0 +1,466 @@ +package shadowsocks_test + +import ( + "bytes" + "errors" + "net" + "testing" + + "github.com/33TU/socks/shadowsocks" +) + +func TestTCPRequestFixedHeader_Init_Validate(t *testing.T) { + tests := []struct { + name string + hdr shadowsocks.TCPRequestFixedHeader + wantErr error + }{ + { + name: "valid", + hdr: func() shadowsocks.TCPRequestFixedHeader { + var h shadowsocks.TCPRequestFixedHeader + h.Init(shadowsocks.TCPHeaderTypeClientStream, 123456789, 42) + return h + }(), + }, + { + name: "invalid type", + hdr: func() shadowsocks.TCPRequestFixedHeader { + var h shadowsocks.TCPRequestFixedHeader + h.Init(0x99, 123456789, 42) + return h + }(), + wantErr: shadowsocks.ErrInvalidTCPHeaderType, + }, + } + + for _, tt := range tests { + t.Run(tt.name, func(t *testing.T) { + err := tt.hdr.Validate() + if !errors.Is(err, tt.wantErr) { + t.Fatalf("Validate() error = %v, wantErr = %v", err, tt.wantErr) + } + }) + } +} + +func TestTCPRequestFixedHeader_EncodedLen(t *testing.T) { + var h shadowsocks.TCPRequestFixedHeader + h.Init(shadowsocks.TCPHeaderTypeClientStream, 1, 2) + + if got := h.EncodedLen(); got != shadowsocks.TcpRequestFixedHeaderLen { + t.Fatalf("EncodedLen() = %d, want %d", got, shadowsocks.TcpRequestFixedHeaderLen) + } +} + +func TestTCPRequestFixedHeader_EncodeTo_Decode_RoundTrip(t *testing.T) { + var want shadowsocks.TCPRequestFixedHeader + want.Init(shadowsocks.TCPHeaderTypeClientStream, 123456789, 321) + + buf := make([]byte, want.EncodedLen()) + + bw, err := want.EncodeTo(buf[:0]) + if err != nil { + t.Fatalf("EncodeTo() failed: %v", err) + } + if len(bw) != len(buf) { + t.Fatalf("EncodeTo() wrote %d bytes, want %d", len(bw), len(buf)) + } + + var got shadowsocks.TCPRequestFixedHeader + nr, err := got.Decode(buf) + if err != nil { + t.Fatalf("Decode() failed: %v", err) + } + if nr != len(buf) { + t.Fatalf("Decode() read %d bytes, want %d", nr, len(buf)) + } + + if got.Type != want.Type || got.Timestamp != want.Timestamp || got.Length != want.Length { + t.Fatalf("round-trip mismatch: got %+v, want %+v", got, want) + } +} + +func TestTCPRequestFixedHeader_EncodeTo_Invalid(t *testing.T) { + tests := []struct { + name string + hdr shadowsocks.TCPRequestFixedHeader + bufLen int + wantErr error + }{ + { + name: "invalid type", + hdr: shadowsocks.TCPRequestFixedHeader{ + Type: 0x99, + Timestamp: 1, + Length: 2, + }, + bufLen: 32, + wantErr: shadowsocks.ErrInvalidTCPHeaderType, + }, + } + + for _, tt := range tests { + t.Run(tt.name, func(t *testing.T) { + buf := make([]byte, tt.bufLen) + _, err := tt.hdr.EncodeTo(buf) + if !errors.Is(err, tt.wantErr) { + t.Fatalf("EncodeTo() error = %v, wantErr = %v", err, tt.wantErr) + } + }) + } +} + +func TestTCPRequestFixedHeader_Decode_Invalid(t *testing.T) { + tests := []struct { + name string + src []byte + wantErr error + }{ + { + name: "short header", + src: make([]byte, shadowsocks.TcpRequestFixedHeaderLen-1), + wantErr: shadowsocks.ErrShortTCPHeader, + }, + { + name: "invalid type", + src: []byte{ + 0x99, + 0, 0, 0, 0, 0, 0, 0, 1, + 0, 2, + }, + wantErr: shadowsocks.ErrInvalidTCPHeaderType, + }, + } + + for _, tt := range tests { + t.Run(tt.name, func(t *testing.T) { + var h shadowsocks.TCPRequestFixedHeader + _, err := h.Decode(tt.src) + if !errors.Is(err, tt.wantErr) { + t.Fatalf("Decode() error = %v, wantErr = %v", err, tt.wantErr) + } + }) + } +} + +func TestTCPRequestFixedHeader_String(t *testing.T) { + var h shadowsocks.TCPRequestFixedHeader + h.Init(shadowsocks.TCPHeaderTypeClientStream, 123, 45) + + want := "TCPRequestFixedHeader{Type:0 Timestamp:123 Length:45}" + if got := h.String(); got != want { + t.Fatalf("String() = %q, want %q", got, want) + } +} + +/////// + +func TestTCPRequestVariableHeader_Init_Validate(t *testing.T) { + validTarget := shadowsocks.Addr{ + AddrType: shadowsocks.AddrTypeDomain, + Domain: "example.com", + Port: 443, + } + + tests := []struct { + name string + hdr shadowsocks.TCPRequestVariableHeader + wantErr error + }{ + { + name: "valid with padding", + hdr: func() shadowsocks.TCPRequestVariableHeader { + var h shadowsocks.TCPRequestVariableHeader + h.Init(validTarget, []byte{1, 2, 3}, nil) + return h + }(), + }, + { + name: "valid with initial data", + hdr: func() shadowsocks.TCPRequestVariableHeader { + var h shadowsocks.TCPRequestVariableHeader + h.Init(validTarget, nil, []byte("hello")) + return h + }(), + }, + { + name: "invalid target", + hdr: shadowsocks.TCPRequestVariableHeader{ + Target: shadowsocks.Addr{ + AddrType: 0x99, + Port: 80, + }, + PaddingLen: 1, + Padding: []byte{1}, + }, + wantErr: shadowsocks.ErrInvalidAddrType, + }, + { + name: "invalid padding length", + hdr: shadowsocks.TCPRequestVariableHeader{ + Target: validTarget, + PaddingLen: 5, + Padding: []byte{1, 2}, + }, + wantErr: shadowsocks.ErrInvalidTCPPaddingLength, + }, + { + name: "missing header data", + hdr: shadowsocks.TCPRequestVariableHeader{ + Target: validTarget, + PaddingLen: 0, + Padding: nil, + InitialData: nil, + }, + wantErr: shadowsocks.ErrMissingTCPHeaderData, + }, + } + + for _, tt := range tests { + t.Run(tt.name, func(t *testing.T) { + err := tt.hdr.Validate() + if !errors.Is(err, tt.wantErr) { + t.Fatalf("Validate() error = %v, wantErr = %v", err, tt.wantErr) + } + }) + } +} + +func TestTCPRequestVariableHeader_EncodedLen(t *testing.T) { + h := shadowsocks.TCPRequestVariableHeader{ + Target: shadowsocks.Addr{ + AddrType: shadowsocks.AddrTypeDomain, + Domain: "example.com", + Port: 443, + }, + PaddingLen: 3, + Padding: []byte{1, 2, 3}, + InitialData: []byte("abc"), + } + + want := h.Target.EncodedLen() + 2 + 3 + 3 + if got := h.EncodedLen(); got != want { + t.Fatalf("EncodedLen() = %d, want %d", got, want) + } +} + +func TestTCPRequestVariableHeader_EncodeTo_Decode_RoundTrip(t *testing.T) { + tests := []struct { + name string + hdr shadowsocks.TCPRequestVariableHeader + }{ + { + name: "domain target", + hdr: shadowsocks.TCPRequestVariableHeader{ + Target: shadowsocks.Addr{ + AddrType: shadowsocks.AddrTypeDomain, + Domain: "example.com", + Port: 443, + }, + PaddingLen: 2, + Padding: []byte{0xaa, 0xbb}, + InitialData: []byte("hello"), + }, + }, + { + name: "ipv4 target", + hdr: shadowsocks.TCPRequestVariableHeader{ + Target: shadowsocks.Addr{ + AddrType: shadowsocks.AddrTypeIPv4, + IP: net.IPv4(127, 0, 0, 1), + Port: 1080, + }, + PaddingLen: 1, + Padding: []byte{0x01}, + InitialData: []byte("x"), + }, + }, + { + name: "ipv6 target", + hdr: shadowsocks.TCPRequestVariableHeader{ + Target: shadowsocks.Addr{ + AddrType: shadowsocks.AddrTypeIPv6, + IP: net.ParseIP("2001:db8::1"), + Port: 53, + }, + PaddingLen: 3, + Padding: []byte{1, 2, 3}, + InitialData: []byte("payload"), + }, + }, + } + + for _, tt := range tests { + t.Run(tt.name, func(t *testing.T) { + buf := make([]byte, tt.hdr.EncodedLen()) + + bw, err := tt.hdr.EncodeTo(buf[:0]) + if err != nil { + t.Fatalf("EncodeTo() failed: %v", err) + } + if len(bw) != len(buf) { + t.Fatalf("EncodeTo() wrote %d bytes, want %d", len(bw), len(buf)) + } + + var got shadowsocks.TCPRequestVariableHeader + nr, err := got.Decode(buf) + if err != nil { + t.Fatalf("Decode() failed: %v", err) + } + if nr != len(buf) { + t.Fatalf("Decode() read %d bytes, want %d", nr, len(buf)) + } + + if got.Target.AddrType != tt.hdr.Target.AddrType { + t.Fatalf("Target.AddrType = %v, want %v", got.Target.AddrType, tt.hdr.Target.AddrType) + } + if got.Target.Port != tt.hdr.Target.Port { + t.Fatalf("Target.Port = %d, want %d", got.Target.Port, tt.hdr.Target.Port) + } + if got.Target.Domain != tt.hdr.Target.Domain { + t.Fatalf("Target.Domain = %q, want %q", got.Target.Domain, tt.hdr.Target.Domain) + } + if tt.hdr.Target.IP != nil && !got.Target.IP.Equal(tt.hdr.Target.IP) { + t.Fatalf("Target.IP = %v, want %v", got.Target.IP, tt.hdr.Target.IP) + } + if got.PaddingLen != tt.hdr.PaddingLen { + t.Fatalf("PaddingLen = %d, want %d", got.PaddingLen, tt.hdr.PaddingLen) + } + if !bytes.Equal(got.Padding, tt.hdr.Padding) { + t.Fatalf("Padding = %v, want %v", got.Padding, tt.hdr.Padding) + } + if !bytes.Equal(got.InitialData, tt.hdr.InitialData) { + t.Fatalf("InitialData = %v, want %v", got.InitialData, tt.hdr.InitialData) + } + }) + } +} + +func TestTCPRequestVariableHeader_EncodeTo_Invalid(t *testing.T) { + validTarget := shadowsocks.Addr{ + AddrType: shadowsocks.AddrTypeDomain, + Domain: "example.com", + Port: 443, + } + + tests := []struct { + name string + hdr shadowsocks.TCPRequestVariableHeader + bufLen int + wantErr error + }{ + { + name: "invalid target", + hdr: shadowsocks.TCPRequestVariableHeader{ + Target: shadowsocks.Addr{ + AddrType: 0x99, + }, + PaddingLen: 1, + Padding: []byte{1}, + }, + bufLen: 64, + wantErr: shadowsocks.ErrInvalidAddrType, + }, + { + name: "invalid padding length", + hdr: shadowsocks.TCPRequestVariableHeader{ + Target: validTarget, + PaddingLen: 10, + Padding: []byte{1}, + }, + bufLen: 64, + wantErr: shadowsocks.ErrInvalidTCPPaddingLength, + }, + { + name: "missing header data", + hdr: shadowsocks.TCPRequestVariableHeader{ + Target: validTarget, + }, + bufLen: 64, + wantErr: shadowsocks.ErrMissingTCPHeaderData, + }, + } + + for _, tt := range tests { + t.Run(tt.name, func(t *testing.T) { + buf := make([]byte, tt.bufLen) + _, err := tt.hdr.EncodeTo(buf) + if !errors.Is(err, tt.wantErr) { + t.Fatalf("EncodeTo() error = %v, wantErr = %v", err, tt.wantErr) + } + }) + } +} + +func TestTCPRequestVariableHeader_Decode_Invalid(t *testing.T) { + validTarget := shadowsocks.Addr{ + AddrType: shadowsocks.AddrTypeDomain, + Domain: "example.com", + Port: 443, + } + validTargetBuf, err := validTarget.EncodeTo(nil) + if err != nil { + t.Fatalf("target EncodeTo() failed: %v", err) + } + + tests := []struct { + name string + src []byte + wantErr error + }{ + { + name: "invalid target", + src: []byte{0x99}, + wantErr: shadowsocks.ErrInvalidAddrType, + }, + { + name: "short after target before padding length", + src: validTargetBuf[:len(validTargetBuf)-1], + wantErr: shadowsocks.ErrShortAddr, + }, + { + name: "missing padding length bytes", + src: append(append([]byte{}, validTargetBuf...), 0x00), + wantErr: shadowsocks.ErrShortTCPHeader, + }, + { + name: "short padding bytes", + src: append(append(append([]byte{}, validTargetBuf...), 0x00, 0x02), 0x01), + wantErr: shadowsocks.ErrShortTCPHeader, + }, + { + name: "missing header data", + src: append(append([]byte{}, validTargetBuf...), 0x00, 0x00), + wantErr: shadowsocks.ErrMissingTCPHeaderData, + }, + } + + for _, tt := range tests { + t.Run(tt.name, func(t *testing.T) { + var h shadowsocks.TCPRequestVariableHeader + _, err := h.Decode(tt.src) + if !errors.Is(err, tt.wantErr) { + t.Fatalf("Decode() error = %v, wantErr = %v", err, tt.wantErr) + } + }) + } +} + +func TestTCPRequestVariableHeader_String(t *testing.T) { + h := shadowsocks.TCPRequestVariableHeader{ + Target: shadowsocks.Addr{ + AddrType: shadowsocks.AddrTypeDomain, + Domain: "example.com", + Port: 443, + }, + PaddingLen: 2, + Padding: []byte{1, 2}, + InitialData: []byte("abc"), + } + + want := `TCPRequestVariableHeader{Target:Addr{AddrType=DOMAIN, Host=example.com, Port=443} PaddingLen:2 InitialDataLen:3}` + if got := h.String(); got != want { + t.Fatalf("String() = %q, want %q", got, want) + } +} diff --git a/shadowsocks/tcp_response_header.go b/shadowsocks/tcp_response_header.go new file mode 100644 index 0000000..408d752 --- /dev/null +++ b/shadowsocks/tcp_response_header.go @@ -0,0 +1,85 @@ +package shadowsocks + +import ( + "encoding/binary" + "fmt" +) + +// TCPResponseHeader represents the fixed-length response header used by +// Shadowsocks 2022 TCP streams. +type TCPResponseHeader struct { + Type byte + Timestamp uint64 + RequestSalt []byte + Length uint16 +} + +// Init initializes a TCPResponseHeader. +func (h *TCPResponseHeader) Init(typ byte, timestamp uint64, requestSalt []byte, length uint16) { + h.Type = typ + h.Timestamp = timestamp + h.RequestSalt = requestSalt + h.Length = length +} + +// Validate checks the correctness of the response header fields. +func (h *TCPResponseHeader) Validate() error { + if h.Type != TCPHeaderTypeServerStream { + return ErrInvalidTCPHeaderType + } + if len(h.RequestSalt) == 0 { + return ErrMissingTCPResponseSalt + } + return nil +} + +// EncodedLen returns the number of bytes required to encode the response header. +func (h *TCPResponseHeader) EncodedLen() int { + return 1 + 8 + len(h.RequestSalt) + 2 +} + +// Decode decodes a response header from src using the expected request salt length. +// It returns the number of bytes consumed. +func (h *TCPResponseHeader) Decode(src []byte, requestSaltLen int) (int, error) { + if requestSaltLen <= 0 { + return 0, ErrInvalidTCPResponseSaltLen + } + + need := 1 + 8 + requestSaltLen + 2 + if len(src) < need { + return 0, ErrShortTCPHeader + } + + h.Type = src[0] + h.Timestamp = binary.BigEndian.Uint64(src[1:9]) + h.RequestSalt = append(h.RequestSalt[:0], src[9:9+requestSaltLen]...) + h.Length = binary.BigEndian.Uint16(src[9+requestSaltLen : 9+requestSaltLen+2]) + + if err := h.Validate(); err != nil { + return 0, err + } + + return need, nil +} + +// EncodeTo encodes the response header into dst and returns the extended slice. +func (h *TCPResponseHeader) EncodeTo(dst []byte) ([]byte, error) { + if err := h.Validate(); err != nil { + return nil, err + } + + dst = append(dst, h.Type) + dst = binary.BigEndian.AppendUint64(dst, h.Timestamp) + dst = append(dst, h.RequestSalt...) + dst = binary.BigEndian.AppendUint16(dst, h.Length) + + return dst, nil +} + +// String returns a human-readable representation of the response header. +func (h *TCPResponseHeader) String() string { + return fmt.Sprintf( + "TCPResponseHeader{Type:%d Timestamp:%d RequestSaltLen:%d Length:%d}", + h.Type, h.Timestamp, len(h.RequestSalt), h.Length, + ) +} diff --git a/shadowsocks/tcp_response_header_test.go b/shadowsocks/tcp_response_header_test.go new file mode 100644 index 0000000..4809e08 --- /dev/null +++ b/shadowsocks/tcp_response_header_test.go @@ -0,0 +1,220 @@ +package shadowsocks_test + +import ( + "bytes" + "errors" + "testing" + + "github.com/33TU/socks/shadowsocks" +) + +func TestTCPResponseHeader_Init_Validate(t *testing.T) { + tests := []struct { + name string + hdr shadowsocks.TCPResponseHeader + wantErr error + }{ + { + name: "valid", + hdr: func() shadowsocks.TCPResponseHeader { + var h shadowsocks.TCPResponseHeader + h.Init(shadowsocks.TCPHeaderTypeServerStream, 123456789, []byte{1, 2, 3, 4}, 4) + return h + }(), + }, + { + name: "invalid type", + hdr: func() shadowsocks.TCPResponseHeader { + var h shadowsocks.TCPResponseHeader + h.Init(0x99, 123456789, []byte{1, 2, 3, 4}, 4) + return h + }(), + wantErr: shadowsocks.ErrInvalidTCPHeaderType, + }, + { + name: "missing salt", + hdr: func() shadowsocks.TCPResponseHeader { + var h shadowsocks.TCPResponseHeader + h.Init(shadowsocks.TCPHeaderTypeServerStream, 123456789, nil, 0) + return h + }(), + wantErr: shadowsocks.ErrMissingTCPResponseSalt, + }, + } + + for _, tt := range tests { + t.Run(tt.name, func(t *testing.T) { + err := tt.hdr.Validate() + if !errors.Is(err, tt.wantErr) { + t.Fatalf("Validate() error = %v, wantErr = %v", err, tt.wantErr) + } + }) + } +} + +func TestTCPResponseHeader_EncodedLen(t *testing.T) { + h := shadowsocks.TCPResponseHeader{ + Type: shadowsocks.TCPHeaderTypeServerStream, + Timestamp: 1, + RequestSalt: []byte{1, 2, 3, 4}, + Length: 4, + } + + want := shadowsocks.TcpResponseFixedBaseLen + 4 + if got := h.EncodedLen(); got != want { + t.Fatalf("EncodedLen() = %d, want %d", got, want) + } +} + +func TestTCPResponseHeader_EncodeTo_Decode_RoundTrip(t *testing.T) { + want := shadowsocks.TCPResponseHeader{ + Type: shadowsocks.TCPHeaderTypeServerStream, + Timestamp: 123456789, + RequestSalt: []byte{0xaa, 0xbb, 0xcc, 0xdd}, + Length: 4, // first response payload length + } + + buf, err := want.EncodeTo(nil) + if err != nil { + t.Fatalf("EncodeTo() failed: %v", err) + } + if len(buf) != want.EncodedLen() { + t.Fatalf("EncodeTo() wrote %d bytes, want %d", len(buf), want.EncodedLen()) + } + + var got shadowsocks.TCPResponseHeader + nr, err := got.Decode(buf, len(want.RequestSalt)) + if err != nil { + t.Fatalf("Decode() failed: %v", err) + } + if nr != len(buf) { + t.Fatalf("Decode() read %d bytes, want %d", nr, len(buf)) + } + + if got.Type != want.Type || got.Timestamp != want.Timestamp || got.Length != want.Length { + t.Fatalf("header mismatch: got %+v, want %+v", got, want) + } + if !bytes.Equal(got.RequestSalt, want.RequestSalt) { + t.Fatalf("RequestSalt = %v, want %v", got.RequestSalt, want.RequestSalt) + } +} + +func TestTCPResponseHeader_EncodeTo_Invalid(t *testing.T) { + tests := []struct { + name string + hdr shadowsocks.TCPResponseHeader + bufLen int + wantErr error + }{ + { + name: "invalid type", + hdr: shadowsocks.TCPResponseHeader{ + Type: 0x99, + Timestamp: 1, + RequestSalt: []byte{1, 2}, + Length: 2, + }, + bufLen: 32, + wantErr: shadowsocks.ErrInvalidTCPHeaderType, + }, + { + name: "missing salt", + hdr: shadowsocks.TCPResponseHeader{ + Type: shadowsocks.TCPHeaderTypeServerStream, + Timestamp: 1, + Length: 0, + }, + bufLen: 32, + wantErr: shadowsocks.ErrMissingTCPResponseSalt, + }, + } + + for _, tt := range tests { + t.Run(tt.name, func(t *testing.T) { + buf := make([]byte, tt.bufLen) + _, err := tt.hdr.EncodeTo(buf[:0]) + if !errors.Is(err, tt.wantErr) { + t.Fatalf("EncodeTo() error = %v, wantErr = %v", err, tt.wantErr) + } + }) + } +} + +func TestTCPResponseHeader_Decode_Invalid(t *testing.T) { + tests := []struct { + name string + src []byte + requestSaltLen int + wantErr error + }{ + { + name: "invalid request salt length", + src: []byte{}, + requestSaltLen: 0, + wantErr: shadowsocks.ErrInvalidTCPResponseSaltLen, + }, + { + name: "short fixed header", + src: make([]byte, 1+8+2+2-1), // type + timestamp + 2-byte salt + length - 1 + requestSaltLen: 2, + wantErr: shadowsocks.ErrShortTCPHeader, + }, + { + name: "invalid type", + src: []byte{ + 0x99, + 0, 0, 0, 0, 0, 0, 0, 1, + 0xaa, 0xbb, + 0, 2, + }, + requestSaltLen: 2, + wantErr: shadowsocks.ErrInvalidTCPHeaderType, + }, + { + name: "missing salt bytes", + src: []byte{ + shadowsocks.TCPHeaderTypeServerStream, + 0, 0, 0, 0, 0, 0, 0, 1, + 0xaa, + 0, 2, + }, + requestSaltLen: 2, + wantErr: shadowsocks.ErrShortTCPHeader, + }, + { + name: "missing length bytes after salt", + src: []byte{ + shadowsocks.TCPHeaderTypeServerStream, + 0, 0, 0, 0, 0, 0, 0, 1, + 0xaa, 0xbb, + 0, + }, + requestSaltLen: 2, + wantErr: shadowsocks.ErrShortTCPHeader, + }, + } + + for _, tt := range tests { + t.Run(tt.name, func(t *testing.T) { + var h shadowsocks.TCPResponseHeader + _, err := h.Decode(tt.src, tt.requestSaltLen) + if !errors.Is(err, tt.wantErr) { + t.Fatalf("Decode() error = %v, wantErr = %v", err, tt.wantErr) + } + }) + } +} + +func TestTCPResponseHeader_String(t *testing.T) { + h := shadowsocks.TCPResponseHeader{ + Type: shadowsocks.TCPHeaderTypeServerStream, + Timestamp: 123, + RequestSalt: []byte{1, 2, 3, 4}, + Length: 4, + } + + want := "TCPResponseHeader{Type:1 Timestamp:123 RequestSaltLen:4 Length:4}" + if got := h.String(); got != want { + t.Fatalf("String() = %q, want %q", got, want) + } +} diff --git a/shadowsocks/tcp_start.go b/shadowsocks/tcp_start.go new file mode 100644 index 0000000..5c77682 --- /dev/null +++ b/shadowsocks/tcp_start.go @@ -0,0 +1,318 @@ +package shadowsocks + +import ( + "bytes" + "fmt" + "io" + "time" + + "github.com/33TU/socks/internal" +) + +const ( + tcpRequestStartStackBufSize = 1024 + tcpResponseStartStackBufSize = 256 +) + +// ParsedTCPRequestStart is the parsed startup state for a TCP request stream. +type ParsedTCPRequestStart struct { + Salt []byte + Cipher *TCPStreamCipher + Fixed TCPRequestFixedHeader + Header TCPRequestVariableHeader +} + +func (s *ParsedTCPRequestStart) Validate(method Method) error { + if s == nil { + return fmt.Errorf("nil parsed TCP request start") + } + if err := method.Validate(); err != nil { + return err + } + if len(s.Salt) != method.SaltSize { + return fmt.Errorf("invalid request salt length: got %d, want %d", len(s.Salt), method.SaltSize) + } + if s.Cipher == nil { + return fmt.Errorf("missing request cipher") + } + if err := s.Cipher.Validate(); err != nil { + return err + } + if err := s.Fixed.Validate(); err != nil { + return err + } + if err := s.Header.Validate(); err != nil { + return err + } + if int(s.Fixed.Length) != s.Header.EncodedLen() { + return fmt.Errorf("request variable header length mismatch: got %d, want %d", s.Fixed.Length, s.Header.EncodedLen()) + } + return nil +} + +// ParsedTCPResponseStart is the parsed startup state for a TCP response stream. +type ParsedTCPResponseStart struct { + Salt []byte + Cipher *TCPStreamCipher + Header TCPResponseHeader + InitialPayload []byte +} + +func (s *ParsedTCPResponseStart) Validate(method Method, expectedRequestSalt []byte) error { + if s == nil { + return fmt.Errorf("nil parsed TCP response start") + } + if err := method.Validate(); err != nil { + return err + } + if len(s.Salt) != method.SaltSize { + return fmt.Errorf("invalid response salt length: got %d, want %d", len(s.Salt), method.SaltSize) + } + if s.Cipher == nil { + return fmt.Errorf("missing response cipher") + } + if err := s.Cipher.Validate(); err != nil { + return err + } + if err := s.Header.Validate(); err != nil { + return err + } + if len(s.Header.RequestSalt) != len(expectedRequestSalt) { + return fmt.Errorf("invalid echoed request salt length: got %d, want %d", len(s.Header.RequestSalt), len(expectedRequestSalt)) + } + if !bytes.Equal(s.Header.RequestSalt, expectedRequestSalt) { + return fmt.Errorf("response request salt mismatch") + } + if int(s.Header.Length) != len(s.InitialPayload) { + return fmt.Errorf("invalid initial payload length: got %d, want %d", len(s.InitialPayload), s.Header.Length) + } + return nil +} + +func EncodedTCPRequestStartLen(method Method, requestSalt []byte, variableHeaderLen int) (int, error) { + if err := method.Validate(); err != nil { + return 0, err + } + if len(requestSalt) != method.SaltSize { + return 0, fmt.Errorf("invalid request salt length: got %d, want %d", len(requestSalt), method.SaltSize) + } + if variableHeaderLen < 0 { + return 0, fmt.Errorf("invalid variable header length: %d", variableHeaderLen) + } + return len(requestSalt) + (TcpRequestFixedHeaderLen + method.TagSize) + (variableHeaderLen + method.TagSize), nil +} + +// WriteTCPRequestStart writes request salt, encrypted fixed header, and encrypted variable header. +func WriteTCPRequestStart(dst io.Writer, method Method, psk, requestSalt []byte, timestamp time.Time, target Addr, padding, initialData []byte) (*TCPStreamCipher, int64, error) { + if err := method.Validate(); err != nil { + return nil, 0, err + } + requestCipher, err := NewTCPStreamCipherFromPSK(method, psk, requestSalt) + if err != nil { + return nil, 0, err + } + + var variableHeader TCPRequestVariableHeader + variableHeader.Init(target, padding, initialData) + if err := variableHeader.Validate(); err != nil { + return nil, 0, err + } + + var fixedHeader TCPRequestFixedHeader + fixedHeader.Init(TCPHeaderTypeClientStream, uint64(timestamp.Unix()), uint16(variableHeader.EncodedLen())) + + scratchLen := TcpRequestFixedHeaderLen + if variableHeader.EncodedLen() > scratchLen { + scratchLen = variableHeader.EncodedLen() + } + plainScratch := internal.GetBytes(scratchLen) + defer internal.PutBytes(plainScratch) + + var stackBuf [tcpRequestStartStackBufSize]byte + out := stackBuf[:0] + out = append(out, requestSalt...) + + out, err = requestCipher.EncodeRequestFixedHeaderTo(out, &fixedHeader, plainScratch[:0]) + if err != nil { + return nil, 0, err + } + out, err = requestCipher.EncodeRequestVariableHeaderTo(out, &variableHeader, plainScratch[:0]) + if err != nil { + return nil, 0, err + } + + n, err := dst.Write(out) + return requestCipher, int64(n), err +} + +// ReadTCPRequestStart reads and decrypts a full client request startup. +func ReadTCPRequestStart(src io.Reader, method Method, psk []byte) (*ParsedTCPRequestStart, int64, error) { + var total int64 + if err := method.Validate(); err != nil { + return nil, 0, err + } + if len(psk) != method.KeySize { + return nil, 0, fmt.Errorf("invalid PSK length: got %d, want %d", len(psk), method.KeySize) + } + + requestSaltBuf := internal.GetBytes(method.SaltSize) + defer internal.PutBytes(requestSaltBuf) + n, err := io.ReadFull(src, requestSaltBuf) + total += int64(n) + if err != nil { + return nil, total, err + } + + requestCipher, err := NewTCPStreamCipherFromPSK(method, psk, requestSaltBuf) + if err != nil { + return nil, total, err + } + + encFixed := internal.GetBytes(TcpRequestFixedHeaderLen + method.TagSize) + defer internal.PutBytes(encFixed) + n, err = io.ReadFull(src, encFixed) + total += int64(n) + if err != nil { + return nil, total, err + } + + fixedPlainScratch := internal.GetBytes(TcpRequestFixedHeaderLen) + defer internal.PutBytes(fixedPlainScratch) + fixedHeader, err := requestCipher.DecodeRequestFixedHeader(encFixed, fixedPlainScratch[:0]) + if err != nil { + return nil, total, err + } + + encVariable := internal.GetBytes(int(fixedHeader.Length) + method.TagSize) + defer internal.PutBytes(encVariable) + n, err = io.ReadFull(src, encVariable) + total += int64(n) + if err != nil { + return nil, total, err + } + + variablePlainScratch := internal.GetBytes(int(fixedHeader.Length)) + defer internal.PutBytes(variablePlainScratch) + variableHeader, err := requestCipher.DecodeRequestVariableHeader(encVariable, variablePlainScratch[:0]) + if err != nil { + return nil, total, err + } + + parsed := &ParsedTCPRequestStart{ + Salt: append([]byte(nil), requestSaltBuf...), + Cipher: requestCipher, + Fixed: fixedHeader, + Header: variableHeader, + } + if err := parsed.Validate(method); err != nil { + return nil, total, err + } + return parsed, total, nil +} + +// WriteTCPResponseStart writes response salt, encrypted response header, and first encrypted payload. +func WriteTCPResponseStart(dst io.Writer, method Method, psk, responseSalt []byte, timestamp time.Time, requestSalt, initialPayload []byte) (*TCPStreamCipher, int64, error) { + if err := method.Validate(); err != nil { + return nil, 0, err + } + if len(requestSalt) != method.SaltSize { + return nil, 0, fmt.Errorf("invalid request salt length: got %d, want %d", len(requestSalt), method.SaltSize) + } + responseCipher, err := NewTCPStreamCipherFromPSK(method, psk, responseSalt) + if err != nil { + return nil, 0, err + } + + var header TCPResponseHeader + header.Init(TCPHeaderTypeServerStream, uint64(timestamp.Unix()), requestSalt, uint16(len(initialPayload))) + if err := header.Validate(); err != nil { + return nil, 0, err + } + + headerPlainScratch := internal.GetBytes(header.EncodedLen()) + defer internal.PutBytes(headerPlainScratch) + + var stackBuf [tcpResponseStartStackBufSize]byte + out := stackBuf[:0] + out = append(out, responseSalt...) + out, err = responseCipher.EncodeResponseHeaderTo(out, &header, headerPlainScratch[:0]) + if err != nil { + return nil, 0, err + } + out, err = responseCipher.EncodeChunkPayloadTo(out, initialPayload) + if err != nil { + return nil, 0, err + } + + n, err := dst.Write(out) + return responseCipher, int64(n), err +} + +// ReadTCPResponseStart reads and decrypts a full server response startup. +func ReadTCPResponseStart(src io.Reader, method Method, psk, expectedRequestSalt []byte) (*ParsedTCPResponseStart, int64, error) { + var total int64 + if err := method.Validate(); err != nil { + return nil, 0, err + } + if len(psk) != method.KeySize { + return nil, 0, fmt.Errorf("invalid PSK length: got %d, want %d", len(psk), method.KeySize) + } + if len(expectedRequestSalt) != method.SaltSize { + return nil, 0, fmt.Errorf("invalid request salt length: got %d, want %d", len(expectedRequestSalt), method.SaltSize) + } + + responseSaltBuf := internal.GetBytes(method.SaltSize) + defer internal.PutBytes(responseSaltBuf) + n, err := io.ReadFull(src, responseSaltBuf) + total += int64(n) + if err != nil { + return nil, total, err + } + + responseCipher, err := NewTCPStreamCipherFromPSK(method, psk, responseSaltBuf) + if err != nil { + return nil, total, err + } + + encHeaderLen := 1 + 8 + method.SaltSize + 2 + method.TagSize + encHeader := internal.GetBytes(encHeaderLen) + defer internal.PutBytes(encHeader) + n, err = io.ReadFull(src, encHeader) + total += int64(n) + if err != nil { + return nil, total, err + } + + plainHeaderScratch := internal.GetBytes(1 + 8 + method.SaltSize + 2) + defer internal.PutBytes(plainHeaderScratch) + header, err := responseCipher.DecodeResponseHeader(encHeader, plainHeaderScratch[:0]) + if err != nil { + return nil, total, err + } + + encPayloadBuf := internal.GetBytes(responseCipher.EncryptedPayloadLength(int(header.Length))) + defer internal.PutBytes(encPayloadBuf) + n, err = io.ReadFull(src, encPayloadBuf) + total += int64(n) + if err != nil { + return nil, total, err + } + + payloadScratch := internal.GetBytes(int(header.Length)) + defer internal.PutBytes(payloadScratch) + initialPayload, err := responseCipher.DecodeChunkPayloadTo(payloadScratch[:0], encPayloadBuf) + if err != nil { + return nil, total, err + } + + parsed := &ParsedTCPResponseStart{ + Salt: append([]byte(nil), responseSaltBuf...), + Cipher: responseCipher, + Header: header, + InitialPayload: append([]byte(nil), initialPayload...), + } + if err := parsed.Validate(method, expectedRequestSalt); err != nil { + return nil, total, err + } + return parsed, total, nil +} diff --git a/shadowsocks/tcp_stream_cipher.go b/shadowsocks/tcp_stream_cipher.go new file mode 100644 index 0000000..daf6b71 --- /dev/null +++ b/shadowsocks/tcp_stream_cipher.go @@ -0,0 +1,254 @@ +package shadowsocks + +import ( + "crypto/cipher" + "encoding/binary" + "fmt" +) + +// TCPStreamCipher manages AEAD state for a single Shadowsocks 2022 TCP stream. +type TCPStreamCipher struct { + Method Method + AEAD cipher.AEAD + Nonce [AeadNonceSize]byte +} + +// NewTCPStreamCipher creates a new TCP stream cipher from an already-derived subkey. +func NewTCPStreamCipher(method Method, subkey []byte) (*TCPStreamCipher, error) { + if err := method.Validate(); err != nil { + return nil, err + } + if len(subkey) != method.KeySize { + return nil, fmt.Errorf("invalid subkey length: got %d, want %d", len(subkey), method.KeySize) + } + + aead, err := method.NewAEAD(subkey) + if err != nil { + return nil, err + } + + return &TCPStreamCipher{ + Method: method, + AEAD: aead, + }, nil +} + +// NewTCPStreamCipherFromPSK creates a new TCP stream cipher from a PSK and salt. +func NewTCPStreamCipherFromPSK(method Method, psk, salt []byte) (*TCPStreamCipher, error) { + if err := method.Validate(); err != nil { + return nil, err + } + if len(psk) != method.KeySize { + return nil, fmt.Errorf("invalid PSK length: got %d, want %d", len(psk), method.KeySize) + } + if len(salt) != method.SaltSize { + return nil, fmt.Errorf("invalid salt length: got %d, want %d", len(salt), method.SaltSize) + } + + subkey := make([]byte, method.KeySize) + if err := DeriveSubkeyTo(subkey, method, psk, salt); err != nil { + return nil, err + } + + return NewTCPStreamCipher(method, subkey) +} + +// Reset resets the stream nonce to zero. +func (s *TCPStreamCipher) Reset() { + clear(s.Nonce[:]) +} + +// SealTo encrypts plaintext into dst using the current nonce and increments the nonce. +func (s *TCPStreamCipher) SealTo(dst, plaintext []byte) ([]byte, error) { + if s == nil { + return nil, fmt.Errorf("nil TCP stream cipher") + } + if err := s.Validate(); err != nil { + return nil, err + } + + out := s.AEAD.Seal(dst, s.Nonce[:], plaintext, nil) + s.incNonce() + return out, nil +} + +// OpenTo decrypts ciphertext into dst using the current nonce and increments the nonce. +func (s *TCPStreamCipher) OpenTo(dst, ciphertext []byte) ([]byte, error) { + if s == nil { + return nil, fmt.Errorf("nil TCP stream cipher") + } + if err := s.Validate(); err != nil { + return nil, err + } + + out, err := s.AEAD.Open(dst, s.Nonce[:], ciphertext, nil) + if err != nil { + return nil, err + } + + s.incNonce() + return out, nil +} + +// Validate checks whether the TCP stream cipher is internally valid. +func (s *TCPStreamCipher) Validate() error { + if s == nil { + return fmt.Errorf("nil TCP stream cipher") + } + if err := s.Method.Validate(); err != nil { + return err + } + if s.AEAD == nil { + return fmt.Errorf("missing AEAD") + } + if s.AEAD.NonceSize() != s.Method.NonceSize { + return fmt.Errorf("invalid AEAD nonce size: got %d, want %d", s.AEAD.NonceSize(), s.Method.NonceSize) + } + if s.AEAD.Overhead() != s.Method.TagSize { + return fmt.Errorf("invalid AEAD tag size: got %d, want %d", s.AEAD.Overhead(), s.Method.TagSize) + } + + return nil +} + +// EncryptedChunkLength returns the ciphertext size of a TCP chunk length field. +func (s *TCPStreamCipher) EncryptedChunkLength() int { + return TcpChunkLengthLen + s.Method.TagSize +} + +// EncryptedPayloadLength returns the ciphertext size of a TCP payload chunk. +func (s *TCPStreamCipher) EncryptedPayloadLength(payloadLen int) int { + return payloadLen + s.Method.TagSize +} + +// EncodeChunkLengthTo encrypts a 2-byte big-endian payload length into dst. +func (s *TCPStreamCipher) EncodeChunkLengthTo(dst []byte, payloadLen uint16) ([]byte, error) { + var buf [TcpChunkLengthLen]byte + binary.BigEndian.PutUint16(buf[:], payloadLen) + return s.SealTo(dst, buf[:]) +} + +// DecodeChunkLength decrypts and parses a 2-byte big-endian payload length. +func (s *TCPStreamCipher) DecodeChunkLength(src []byte, scratch []byte) (uint16, error) { + plain, err := s.OpenTo(scratch[:0], src) + if err != nil { + return 0, err + } + if len(plain) != TcpChunkLengthLen { + return 0, fmt.Errorf("invalid TCP chunk length size: got %d, want %d", len(plain), TcpChunkLengthLen) + } + + return binary.BigEndian.Uint16(plain), nil +} + +// EncodeChunkPayloadTo encrypts a TCP payload chunk into dst. +func (s *TCPStreamCipher) EncodeChunkPayloadTo(dst, payload []byte) ([]byte, error) { + if len(payload) > 0xFFFF { + return nil, fmt.Errorf("payload too large: got %d, max %d", len(payload), 0xFFFF) + } + return s.SealTo(dst, payload) +} + +// DecodeChunkPayloadTo decrypts a TCP payload chunk into dst. +func (s *TCPStreamCipher) DecodeChunkPayloadTo(dst, src []byte) ([]byte, error) { + return s.OpenTo(dst, src) +} + +// EncodeRequestFixedHeaderTo encodes and encrypts a TCP request fixed header into dst. +func (s *TCPStreamCipher) EncodeRequestFixedHeaderTo(dst []byte, h *TCPRequestFixedHeader, scratch []byte) ([]byte, error) { + if h == nil { + return nil, fmt.Errorf("nil TCP request fixed header") + } + + plain, err := h.EncodeTo(scratch[:0]) + if err != nil { + return nil, err + } + + return s.SealTo(dst, plain) +} + +// DecodeRequestFixedHeader decrypts and decodes a TCP request fixed header from src. +func (s *TCPStreamCipher) DecodeRequestFixedHeader(src []byte, scratch []byte) (TCPRequestFixedHeader, error) { + var h TCPRequestFixedHeader + + plain, err := s.OpenTo(scratch[:0], src) + if err != nil { + return TCPRequestFixedHeader{}, err + } + if _, err := h.Decode(plain); err != nil { + return TCPRequestFixedHeader{}, err + } + + return h, nil +} + +// EncodeRequestVariableHeaderTo encodes and encrypts a TCP request variable header into dst. +// scratch is used as the plaintext scratch buffer and may be nil. +func (s *TCPStreamCipher) EncodeRequestVariableHeaderTo(dst []byte, h *TCPRequestVariableHeader, scratch []byte) ([]byte, error) { + if h == nil { + return nil, fmt.Errorf("nil TCP request variable header") + } + + plain, err := h.EncodeTo(scratch[:0]) + if err != nil { + return nil, err + } + + return s.SealTo(dst, plain) +} + +// DecodeRequestVariableHeader decrypts and decodes a TCP request variable header from src. +func (s *TCPStreamCipher) DecodeRequestVariableHeader(src []byte, scratch []byte) (TCPRequestVariableHeader, error) { + var h TCPRequestVariableHeader + + plain, err := s.OpenTo(scratch[:0], src) + if err != nil { + return TCPRequestVariableHeader{}, err + } + if _, err := h.Decode(plain); err != nil { + return TCPRequestVariableHeader{}, err + } + + return h, nil +} + +// EncodeResponseHeaderTo encodes and encrypts a TCP response header into dst. +// scratch is used as the plaintext scratch buffer and may be nil. +func (s *TCPStreamCipher) EncodeResponseHeaderTo(dst []byte, h *TCPResponseHeader, scratch []byte) ([]byte, error) { + if h == nil { + return nil, fmt.Errorf("nil TCP response header") + } + + plain, err := h.EncodeTo(scratch[:0]) + if err != nil { + return nil, err + } + + return s.SealTo(dst, plain) +} + +// DecodeResponseHeader decrypts and decodes a TCP response header from src. +func (s *TCPStreamCipher) DecodeResponseHeader(src []byte, scratch []byte) (TCPResponseHeader, error) { + var h TCPResponseHeader + + plain, err := s.OpenTo(scratch[:0], src) + if err != nil { + return TCPResponseHeader{}, err + } + if _, err := h.Decode(plain, s.Method.SaltSize); err != nil { + return TCPResponseHeader{}, err + } + + return h, nil +} + +// incNonce increments the TCP stream nonce as a 12-byte little-endian integer. +func (s *TCPStreamCipher) incNonce() { + for i := range len(s.Nonce) { + s.Nonce[i]++ + if s.Nonce[i] != 0 { + return + } + } +} diff --git a/shadowsocks/tcp_stream_cipher_test.go b/shadowsocks/tcp_stream_cipher_test.go new file mode 100644 index 0000000..8e629cf --- /dev/null +++ b/shadowsocks/tcp_stream_cipher_test.go @@ -0,0 +1,491 @@ +package shadowsocks_test + +import ( + "bytes" + "strings" + "testing" + + "github.com/33TU/socks/shadowsocks" +) + +func newTestMethod(t *testing.T) shadowsocks.Method { + t.Helper() + + method, err := shadowsocks.ParseMethod(shadowsocks.Method2022Blake3AES128GCM) + if err != nil { + t.Fatalf("ParseMethod() error = %v", err) + } + + return method +} + +func newTestPSKAndSalt(method shadowsocks.Method) ([]byte, []byte) { + psk := make([]byte, method.KeySize) + salt := make([]byte, method.SaltSize) + + for i := range psk { + psk[i] = byte(i + 1) + } + for i := range salt { + salt[i] = byte(i + 101) + } + + return psk, salt +} + +func newTestCipherPair(t *testing.T) (*shadowsocks.TCPStreamCipher, *shadowsocks.TCPStreamCipher, shadowsocks.Method) { + t.Helper() + + method := newTestMethod(t) + psk, salt := newTestPSKAndSalt(method) + + enc, err := shadowsocks.NewTCPStreamCipherFromPSK(method, psk, salt) + if err != nil { + t.Fatalf("NewTCPStreamCipherFromPSK(enc) error = %v", err) + } + + dec, err := shadowsocks.NewTCPStreamCipherFromPSK(method, psk, salt) + if err != nil { + t.Fatalf("NewTCPStreamCipherFromPSK(dec) error = %v", err) + } + + return enc, dec, method +} + +func TestNewTCPStreamCipher(t *testing.T) { + t.Parallel() + + method := newTestMethod(t) + + t.Run("valid", func(t *testing.T) { + t.Parallel() + + subkey := make([]byte, method.KeySize) + for i := range subkey { + subkey[i] = byte(i + 1) + } + + s, err := shadowsocks.NewTCPStreamCipher(method, subkey) + if err != nil { + t.Fatalf("NewTCPStreamCipher() error = %v", err) + } + if s == nil { + t.Fatal("NewTCPStreamCipher() returned nil cipher") + } + if err := s.Validate(); err != nil { + t.Fatalf("Validate() error = %v", err) + } + }) + + t.Run("invalid method", func(t *testing.T) { + t.Parallel() + + _, err := shadowsocks.NewTCPStreamCipher(shadowsocks.Method{}, make([]byte, 16)) + if err == nil { + t.Fatal("expected error, got nil") + } + }) + + t.Run("invalid subkey length", func(t *testing.T) { + t.Parallel() + + _, err := shadowsocks.NewTCPStreamCipher(method, make([]byte, method.KeySize-1)) + if err == nil { + t.Fatal("expected error, got nil") + } + if !strings.Contains(err.Error(), "invalid subkey length") { + t.Fatalf("expected invalid subkey length error, got %q", err.Error()) + } + }) +} + +func TestNewTCPStreamCipherFromPSK(t *testing.T) { + t.Parallel() + + method := newTestMethod(t) + psk, salt := newTestPSKAndSalt(method) + + t.Run("valid", func(t *testing.T) { + t.Parallel() + + s, err := shadowsocks.NewTCPStreamCipherFromPSK(method, psk, salt) + if err != nil { + t.Fatalf("NewTCPStreamCipherFromPSK() error = %v", err) + } + if s == nil { + t.Fatal("NewTCPStreamCipherFromPSK() returned nil cipher") + } + if err := s.Validate(); err != nil { + t.Fatalf("Validate() error = %v", err) + } + }) + + t.Run("invalid psk length", func(t *testing.T) { + t.Parallel() + + _, err := shadowsocks.NewTCPStreamCipherFromPSK(method, psk[:len(psk)-1], salt) + if err == nil { + t.Fatal("expected error, got nil") + } + if !strings.Contains(err.Error(), "invalid PSK length") { + t.Fatalf("expected invalid PSK length error, got %q", err.Error()) + } + }) + + t.Run("invalid salt length", func(t *testing.T) { + t.Parallel() + + _, err := shadowsocks.NewTCPStreamCipherFromPSK(method, psk, salt[:len(salt)-1]) + if err == nil { + t.Fatal("expected error, got nil") + } + if !strings.Contains(err.Error(), "invalid salt length") { + t.Fatalf("expected invalid salt length error, got %q", err.Error()) + } + }) +} + +func TestTCPStreamCipher_Reset(t *testing.T) { + t.Parallel() + + s, _, _ := newTestCipherPair(t) + + for i := range s.Nonce { + s.Nonce[i] = 0xff + } + + s.Reset() + + var zero [shadowsocks.AeadNonceSize]byte + if s.Nonce != zero { + t.Fatalf("Nonce = %v, want zero", s.Nonce) + } +} + +func TestTCPStreamCipher_Validate(t *testing.T) { + t.Parallel() + + method := newTestMethod(t) + psk, salt := newTestPSKAndSalt(method) + + valid, err := shadowsocks.NewTCPStreamCipherFromPSK(method, psk, salt) + if err != nil { + t.Fatalf("NewTCPStreamCipherFromPSK() error = %v", err) + } + + t.Run("valid", func(t *testing.T) { + t.Parallel() + + if err := valid.Validate(); err != nil { + t.Fatalf("Validate() error = %v", err) + } + }) + + t.Run("nil cipher", func(t *testing.T) { + t.Parallel() + + var s *shadowsocks.TCPStreamCipher + err := s.Validate() + if err == nil { + t.Fatal("expected error, got nil") + } + if !strings.Contains(err.Error(), "nil TCP stream cipher") { + t.Fatalf("unexpected error: %v", err) + } + }) + + t.Run("missing aead", func(t *testing.T) { + t.Parallel() + + s := &shadowsocks.TCPStreamCipher{ + Method: method, + AEAD: nil, + } + + err := s.Validate() + if err == nil { + t.Fatal("expected error, got nil") + } + if !strings.Contains(err.Error(), "missing AEAD") { + t.Fatalf("unexpected error: %v", err) + } + }) +} + +func TestTCPStreamCipher_SealTo_OpenTo(t *testing.T) { + t.Parallel() + + enc, dec, _ := newTestCipherPair(t) + plaintext := []byte("hello shadowsocks") + + ciphertext, err := enc.SealTo(nil, plaintext) + if err != nil { + t.Fatalf("SealTo() error = %v", err) + } + + got, err := dec.OpenTo(nil, ciphertext) + if err != nil { + t.Fatalf("OpenTo() error = %v", err) + } + + if !bytes.Equal(got, plaintext) { + t.Fatalf("OpenTo() = %q, want %q", got, plaintext) + } +} + +func TestTCPStreamCipher_SealTo_OpenTo_NilCipher(t *testing.T) { + t.Parallel() + + var s *shadowsocks.TCPStreamCipher + + if _, err := s.SealTo(nil, []byte("x")); err == nil { + t.Fatal("expected error from nil SealTo receiver") + } + if _, err := s.OpenTo(nil, []byte("x")); err == nil { + t.Fatal("expected error from nil OpenTo receiver") + } +} + +func TestTCPStreamCipher_OpenTo_WrongCipherFails(t *testing.T) { + t.Parallel() + + enc, _, method := newTestCipherPair(t) + + psk := make([]byte, method.KeySize) + salt := make([]byte, method.SaltSize) + for i := range psk { + psk[i] = byte(i + 9) + } + for i := range salt { + salt[i] = byte(i + 19) + } + + wrongDec, err := shadowsocks.NewTCPStreamCipherFromPSK(method, psk, salt) + if err != nil { + t.Fatalf("NewTCPStreamCipherFromPSK() error = %v", err) + } + + ciphertext, err := enc.SealTo(nil, []byte("hello")) + if err != nil { + t.Fatalf("SealTo() error = %v", err) + } + + if _, err := wrongDec.OpenTo(nil, ciphertext); err == nil { + t.Fatal("expected decrypt error, got nil") + } +} + +func TestTCPStreamCipher_EncryptedLengths(t *testing.T) { + t.Parallel() + + s, _, method := newTestCipherPair(t) + + if got, want := s.EncryptedChunkLength(), shadowsocks.TcpChunkLengthLen+method.TagSize; got != want { + t.Fatalf("EncryptedChunkLength() = %d, want %d", got, want) + } + + if got, want := s.EncryptedPayloadLength(123), 123+method.TagSize; got != want { + t.Fatalf("EncryptedPayloadLength() = %d, want %d", got, want) + } +} + +func TestTCPStreamCipher_ChunkLength_RoundTrip(t *testing.T) { + t.Parallel() + + enc, dec, method := newTestCipherPair(t) + + ciphertext, err := enc.EncodeChunkLengthTo(nil, 1234) + if err != nil { + t.Fatalf("EncodeChunkLengthTo() error = %v", err) + } + + if got, want := len(ciphertext), shadowsocks.TcpChunkLengthLen+method.TagSize; got != want { + t.Fatalf("len(ciphertext) = %d, want %d", got, want) + } + + n, err := dec.DecodeChunkLength(ciphertext, nil) + if err != nil { + t.Fatalf("DecodeChunkLength() error = %v", err) + } + if n != 1234 { + t.Fatalf("DecodeChunkLength() = %d, want %d", n, 1234) + } +} + +func TestTCPStreamCipher_ChunkPayload_RoundTrip(t *testing.T) { + t.Parallel() + + enc, dec, method := newTestCipherPair(t) + payload := []byte("payload-data") + + ciphertext, err := enc.EncodeChunkPayloadTo(nil, payload) + if err != nil { + t.Fatalf("EncodeChunkPayloadTo() error = %v", err) + } + + if got, want := len(ciphertext), len(payload)+method.TagSize; got != want { + t.Fatalf("len(ciphertext) = %d, want %d", got, want) + } + + plain, err := dec.DecodeChunkPayloadTo(nil, ciphertext) + if err != nil { + t.Fatalf("DecodeChunkPayloadTo() error = %v", err) + } + if !bytes.Equal(plain, payload) { + t.Fatalf("DecodeChunkPayloadTo() = %q, want %q", plain, payload) + } +} + +func TestTCPStreamCipher_EncodeChunkPayloadTo_TooLarge(t *testing.T) { + t.Parallel() + + enc, _, _ := newTestCipherPair(t) + payload := make([]byte, 0x10000) + + _, err := enc.EncodeChunkPayloadTo(nil, payload) + if err == nil { + t.Fatal("expected error, got nil") + } + if !strings.Contains(err.Error(), "payload too large") { + t.Fatalf("unexpected error: %v", err) + } +} + +func TestTCPStreamCipher_RequestFixedHeader_RoundTrip(t *testing.T) { + t.Parallel() + + enc, dec, _ := newTestCipherPair(t) + + var want shadowsocks.TCPRequestFixedHeader + want.Init(shadowsocks.TCPHeaderTypeClientStream, 123456789, 321) + + ciphertext, err := enc.EncodeRequestFixedHeaderTo(nil, &want, nil) + if err != nil { + t.Fatalf("EncodeRequestFixedHeaderTo() error = %v", err) + } + + got, err := dec.DecodeRequestFixedHeader(ciphertext, nil) + if err != nil { + t.Fatalf("DecodeRequestFixedHeader() error = %v", err) + } + + if got.Type != want.Type || got.Timestamp != want.Timestamp || got.Length != want.Length { + t.Fatalf("got %+v, want %+v", got, want) + } +} + +func TestTCPStreamCipher_RequestFixedHeader_Nil(t *testing.T) { + t.Parallel() + + enc, _, _ := newTestCipherPair(t) + + _, err := enc.EncodeRequestFixedHeaderTo(nil, nil, nil) + if err == nil { + t.Fatal("expected error, got nil") + } + if !strings.Contains(err.Error(), "nil TCP request fixed header") { + t.Fatalf("unexpected error: %v", err) + } +} + +func TestTCPStreamCipher_RequestVariableHeader_RoundTrip(t *testing.T) { + t.Parallel() + + enc, dec, _ := newTestCipherPair(t) + + var want shadowsocks.TCPRequestVariableHeader + want.Init( + shadowsocks.Addr{ + AddrType: shadowsocks.AddrTypeDomain, + Domain: "example.com", + Port: 443, + }, + []byte{1, 2, 3}, + []byte("hello"), + ) + + ciphertext, err := enc.EncodeRequestVariableHeaderTo(nil, &want, nil) + if err != nil { + t.Fatalf("EncodeRequestVariableHeaderTo() error = %v", err) + } + + got, err := dec.DecodeRequestVariableHeader(ciphertext, nil) + if err != nil { + t.Fatalf("DecodeRequestVariableHeader() error = %v", err) + } + + if got.Target.AddrType != want.Target.AddrType { + t.Fatalf("Target.AddrType = %v, want %v", got.Target.AddrType, want.Target.AddrType) + } + if got.Target.Domain != want.Target.Domain { + t.Fatalf("Target.Domain = %q, want %q", got.Target.Domain, want.Target.Domain) + } + if got.Target.Port != want.Target.Port { + t.Fatalf("Target.Port = %d, want %d", got.Target.Port, want.Target.Port) + } + if got.PaddingLen != want.PaddingLen { + t.Fatalf("PaddingLen = %d, want %d", got.PaddingLen, want.PaddingLen) + } + if !bytes.Equal(got.Padding, want.Padding) { + t.Fatalf("Padding = %v, want %v", got.Padding, want.Padding) + } + if !bytes.Equal(got.InitialData, want.InitialData) { + t.Fatalf("InitialData = %v, want %v", got.InitialData, want.InitialData) + } +} + +func TestTCPStreamCipher_RequestVariableHeader_Nil(t *testing.T) { + t.Parallel() + + enc, _, _ := newTestCipherPair(t) + + _, err := enc.EncodeRequestVariableHeaderTo(nil, nil, nil) + if err == nil { + t.Fatal("expected error, got nil") + } + if !strings.Contains(err.Error(), "nil TCP request variable header") { + t.Fatalf("unexpected error: %v", err) + } +} + +func TestTCPStreamCipher_ResponseHeader_RoundTrip(t *testing.T) { + t.Parallel() + + enc, dec, method := newTestCipherPair(t) + + requestSalt := bytes.Repeat([]byte{0xaa}, method.SaltSize) + + var want shadowsocks.TCPResponseHeader + want.Init(shadowsocks.TCPHeaderTypeServerStream, 123456789, requestSalt, 4) + + ciphertext, err := enc.EncodeResponseHeaderTo(nil, &want, nil) + if err != nil { + t.Fatalf("EncodeResponseHeaderTo() error = %v", err) + } + + got, err := dec.DecodeResponseHeader(ciphertext, nil) + if err != nil { + t.Fatalf("DecodeResponseHeader() error = %v", err) + } + + if got.Type != want.Type || got.Timestamp != want.Timestamp || got.Length != want.Length { + t.Fatalf("got %+v, want %+v", got, want) + } + if !bytes.Equal(got.RequestSalt, want.RequestSalt) { + t.Fatalf("RequestSalt = %v, want %v", got.RequestSalt, want.RequestSalt) + } +} + +func TestTCPStreamCipher_ResponseHeader_Nil(t *testing.T) { + t.Parallel() + + enc, _, _ := newTestCipherPair(t) + + _, err := enc.EncodeResponseHeaderTo(nil, nil, nil) + if err == nil { + t.Fatal("expected error, got nil") + } + if !strings.Contains(err.Error(), "nil TCP response header") { + t.Fatalf("unexpected error: %v", err) + } +}