Skip to content

Commit e0e9153

Browse files
committed
Improved support for SOCK5
1 parent bf36309 commit e0e9153

6 files changed

Lines changed: 822 additions & 19 deletions

File tree

Lines changed: 237 additions & 0 deletions
Original file line numberDiff line numberDiff line change
@@ -0,0 +1,237 @@
1+
package listeners
2+
3+
import (
4+
"bufio"
5+
"errors"
6+
"fmt"
7+
"io"
8+
"net"
9+
"strconv"
10+
)
11+
12+
const (
13+
socks5Version = 0x05
14+
socks5AuthNone = 0x00
15+
socks5AuthUserPass = 0x02
16+
socks5AuthNoMethods = 0xFF
17+
18+
socks5CmdConnect = 0x01
19+
20+
socks5AtypIPv4 = 0x01
21+
socks5AtypDomain = 0x03
22+
socks5AtypIPv6 = 0x04
23+
24+
socks5ReplySucceeded = 0x00
25+
socks5ReplyGeneralFailure = 0x01
26+
socks5ReplyConnectionNotAllow = 0x02
27+
socks5ReplyHostUnreachable = 0x04
28+
socks5ReplyCommandNotSupport = 0x07
29+
socks5ReplyAddressNotSupport = 0x08
30+
)
31+
32+
// SOCKS5AuthConfig controls handshake authentication behavior.
33+
type SOCKS5AuthConfig struct {
34+
Username string
35+
Password string
36+
}
37+
38+
func (c SOCKS5AuthConfig) RequiresUserPass() bool {
39+
return c.Username != "" || c.Password != ""
40+
}
41+
42+
// SOCKS5ConnectRequest is the parsed destination from a CONNECT request.
43+
type SOCKS5ConnectRequest struct {
44+
ATYP byte
45+
Host string
46+
Port int
47+
Target string
48+
}
49+
50+
// PerformSOCKS5Handshake negotiates auth and parses a CONNECT request.
51+
func PerformSOCKS5Handshake(conn net.Conn, auth SOCKS5AuthConfig) (SOCKS5ConnectRequest, error) {
52+
reader := bufio.NewReader(conn)
53+
54+
method, err := negotiateSOCKS5Method(reader, conn, auth)
55+
if err != nil {
56+
return SOCKS5ConnectRequest{}, err
57+
}
58+
if method == socks5AuthUserPass {
59+
if err := authenticateSOCKS5UserPass(reader, conn, auth); err != nil {
60+
return SOCKS5ConnectRequest{}, err
61+
}
62+
}
63+
64+
request, err := readSOCKS5ConnectRequest(reader)
65+
if err != nil {
66+
return SOCKS5ConnectRequest{}, err
67+
}
68+
return request, nil
69+
}
70+
71+
// WriteSOCKS5ConnectReply writes a reply and bind address to the client.
72+
func WriteSOCKS5ConnectReply(conn net.Conn, status byte, bindAddr net.Addr) error {
73+
atyp := byte(socks5AtypIPv4)
74+
addrBytes := []byte{0, 0, 0, 0}
75+
portBytes := []byte{0, 0}
76+
77+
host, port, err := splitHostPort(bindAddr)
78+
if err == nil {
79+
if ip := net.ParseIP(host); ip != nil {
80+
if v4 := ip.To4(); v4 != nil {
81+
atyp = socks5AtypIPv4
82+
addrBytes = v4
83+
} else {
84+
atyp = socks5AtypIPv6
85+
addrBytes = ip.To16()
86+
}
87+
} else {
88+
atyp = socks5AtypDomain
89+
addrBytes = append([]byte{byte(len(host))}, []byte(host)...)
90+
}
91+
portBytes = []byte{byte(port >> 8), byte(port)}
92+
}
93+
94+
reply := []byte{socks5Version, status, 0x00, atyp}
95+
reply = append(reply, addrBytes...)
96+
reply = append(reply, portBytes...)
97+
_, err = conn.Write(reply)
98+
return err
99+
}
100+
101+
func negotiateSOCKS5Method(reader *bufio.Reader, conn net.Conn, auth SOCKS5AuthConfig) (byte, error) {
102+
header := make([]byte, 2)
103+
if _, err := io.ReadFull(reader, header); err != nil {
104+
return 0, err
105+
}
106+
if header[0] != socks5Version {
107+
return 0, fmt.Errorf("unsupported socks version %d", header[0])
108+
}
109+
110+
methods := make([]byte, int(header[1]))
111+
if _, err := io.ReadFull(reader, methods); err != nil {
112+
return 0, err
113+
}
114+
115+
wantMethod := byte(socks5AuthNone)
116+
if auth.RequiresUserPass() {
117+
wantMethod = socks5AuthUserPass
118+
}
119+
120+
for _, method := range methods {
121+
if method == wantMethod {
122+
if _, err := conn.Write([]byte{socks5Version, wantMethod}); err != nil {
123+
return 0, err
124+
}
125+
return wantMethod, nil
126+
}
127+
}
128+
129+
_, _ = conn.Write([]byte{socks5Version, socks5AuthNoMethods})
130+
return 0, errors.New("no compatible socks5 auth method")
131+
}
132+
133+
func authenticateSOCKS5UserPass(reader *bufio.Reader, conn net.Conn, auth SOCKS5AuthConfig) error {
134+
header := make([]byte, 2)
135+
if _, err := io.ReadFull(reader, header); err != nil {
136+
return err
137+
}
138+
if header[0] != 0x01 {
139+
_, _ = conn.Write([]byte{0x01, 0x01})
140+
return fmt.Errorf("unsupported auth version %d", header[0])
141+
}
142+
143+
user := make([]byte, int(header[1]))
144+
if _, err := io.ReadFull(reader, user); err != nil {
145+
return err
146+
}
147+
plen := make([]byte, 1)
148+
if _, err := io.ReadFull(reader, plen); err != nil {
149+
return err
150+
}
151+
pass := make([]byte, int(plen[0]))
152+
if _, err := io.ReadFull(reader, pass); err != nil {
153+
return err
154+
}
155+
156+
if string(user) != auth.Username || string(pass) != auth.Password {
157+
_, _ = conn.Write([]byte{0x01, 0x01})
158+
return errors.New("invalid socks5 username/password")
159+
}
160+
161+
_, err := conn.Write([]byte{0x01, 0x00})
162+
return err
163+
}
164+
165+
func readSOCKS5ConnectRequest(reader *bufio.Reader) (SOCKS5ConnectRequest, error) {
166+
head := make([]byte, 4)
167+
if _, err := io.ReadFull(reader, head); err != nil {
168+
return SOCKS5ConnectRequest{}, err
169+
}
170+
if head[0] != socks5Version {
171+
return SOCKS5ConnectRequest{}, fmt.Errorf("unsupported request version %d", head[0])
172+
}
173+
if head[1] != socks5CmdConnect {
174+
return SOCKS5ConnectRequest{}, fmt.Errorf("unsupported socks5 command %d", head[1])
175+
}
176+
177+
host, err := readSOCKS5Address(reader, head[3])
178+
if err != nil {
179+
return SOCKS5ConnectRequest{}, err
180+
}
181+
portBytes := make([]byte, 2)
182+
if _, err := io.ReadFull(reader, portBytes); err != nil {
183+
return SOCKS5ConnectRequest{}, err
184+
}
185+
port := int(portBytes[0])<<8 | int(portBytes[1])
186+
187+
return SOCKS5ConnectRequest{
188+
ATYP: head[3],
189+
Host: host,
190+
Port: port,
191+
Target: net.JoinHostPort(host, strconv.Itoa(port)),
192+
}, nil
193+
}
194+
195+
func readSOCKS5Address(reader *bufio.Reader, atyp byte) (string, error) {
196+
switch atyp {
197+
case socks5AtypIPv4:
198+
buf := make([]byte, 4)
199+
if _, err := io.ReadFull(reader, buf); err != nil {
200+
return "", err
201+
}
202+
return net.IP(buf).String(), nil
203+
case socks5AtypDomain:
204+
size := make([]byte, 1)
205+
if _, err := io.ReadFull(reader, size); err != nil {
206+
return "", err
207+
}
208+
buf := make([]byte, int(size[0]))
209+
if _, err := io.ReadFull(reader, buf); err != nil {
210+
return "", err
211+
}
212+
return string(buf), nil
213+
case socks5AtypIPv6:
214+
buf := make([]byte, 16)
215+
if _, err := io.ReadFull(reader, buf); err != nil {
216+
return "", err
217+
}
218+
return net.IP(buf).String(), nil
219+
default:
220+
return "", fmt.Errorf("unsupported socks5 atyp %d", atyp)
221+
}
222+
}
223+
224+
func splitHostPort(addr net.Addr) (string, int, error) {
225+
if addr == nil {
226+
return "", 0, errors.New("nil addr")
227+
}
228+
host, portStr, err := net.SplitHostPort(addr.String())
229+
if err != nil {
230+
return "", 0, err
231+
}
232+
port, err := strconv.Atoi(portStr)
233+
if err != nil {
234+
return "", 0, err
235+
}
236+
return host, port, nil
237+
}

internal/dataplane/manager.go

Lines changed: 57 additions & 6 deletions
Original file line numberDiff line numberDiff line change
@@ -29,6 +29,46 @@ func (NoopListenerManager) Start(context.Context) error { return nil }
2929

3030
func (NoopListenerManager) Shutdown(context.Context) error { return nil }
3131

32+
// CompositeListenerManager starts/stops a set of concrete listener managers.
33+
type CompositeListenerManager struct {
34+
managers []ListenerManager
35+
}
36+
37+
func NewCompositeListenerManager(managers ...ListenerManager) *CompositeListenerManager {
38+
active := make([]ListenerManager, 0, len(managers))
39+
for _, manager := range managers {
40+
if manager == nil {
41+
continue
42+
}
43+
active = append(active, manager)
44+
}
45+
return &CompositeListenerManager{managers: active}
46+
}
47+
48+
func (m *CompositeListenerManager) Start(ctx context.Context) error {
49+
started := make([]ListenerManager, 0, len(m.managers))
50+
for _, manager := range m.managers {
51+
if err := manager.Start(ctx); err != nil {
52+
for i := len(started) - 1; i >= 0; i-- {
53+
_ = started[i].Shutdown(ctx)
54+
}
55+
return err
56+
}
57+
started = append(started, manager)
58+
}
59+
return nil
60+
}
61+
62+
func (m *CompositeListenerManager) Shutdown(ctx context.Context) error {
63+
var errs []error
64+
for i := len(m.managers) - 1; i >= 0; i-- {
65+
if err := m.managers[i].Shutdown(ctx); err != nil {
66+
errs = append(errs, err)
67+
}
68+
}
69+
return errors.Join(errs...)
70+
}
71+
3272
// HTTPListenerManager owns one or more HTTP(S) data-plane listeners.
3373
type HTTPListenerManager struct {
3474
drainTimeout time.Duration
@@ -48,27 +88,38 @@ type serverState struct {
4888
conns map[net.Conn]struct{}
4989
}
5090

51-
// NewListenerManager wires the runtime to a concrete listener manager.
52-
// It keeps NoopListenerManager as a fallback when no enabled HTTP listeners exist.
91+
// NewListenerManager wires the runtime to concrete listener managers.
5392
func NewListenerManager(cfg *config.Config) ListenerManager {
5493
if cfg == nil {
5594
return NoopListenerManager{}
5695
}
5796

5897
httpListeners := make([]config.ListenerConfig, 0, len(cfg.Listeners))
98+
socks5Listeners := make([]config.ListenerConfig, 0, len(cfg.Listeners))
5999
for _, listenerCfg := range cfg.Listeners {
60100
if !listenerCfg.Enabled {
61101
continue
62102
}
63-
if listenerCfg.Type == "http" || listenerCfg.Type == "https" {
103+
switch listenerCfg.Type {
104+
case "http", "https":
64105
httpListeners = append(httpListeners, listenerCfg)
106+
case "socks5":
107+
socks5Listeners = append(socks5Listeners, listenerCfg)
65108
}
66109
}
67-
if len(httpListeners) == 0 {
110+
111+
managers := make([]ListenerManager, 0, 2)
112+
runtime := NewRequestRuntime(cfg)
113+
if len(httpListeners) > 0 {
114+
managers = append(managers, NewHTTPListenerManager(httpListeners, defaultDrainTimeout, cfg.Observability.AccessLog.Enabled, runtime))
115+
}
116+
if len(socks5Listeners) > 0 {
117+
managers = append(managers, NewSOCKS5ListenerManager(socks5Listeners, defaultDrainTimeout, runtime))
118+
}
119+
if len(managers) == 0 {
68120
return NoopListenerManager{}
69121
}
70-
71-
return NewHTTPListenerManager(httpListeners, defaultDrainTimeout, cfg.Observability.AccessLog.Enabled, NewRequestRuntime(cfg))
122+
return NewCompositeListenerManager(managers...)
72123
}
73124

74125
func NewHTTPListenerManager(listenerConfigs []config.ListenerConfig, drainTimeout time.Duration, accessLogEnabled bool, runtime listeners.RequestRuntime) *HTTPListenerManager {

internal/dataplane/manager_test.go

Lines changed: 12 additions & 13 deletions
Original file line numberDiff line numberDiff line change
@@ -16,35 +16,34 @@ func TestNewListenerManager_FallbackToNoop(t *testing.T) {
1616
t.Fatalf("expected NoopListenerManager, got %T", mgr)
1717
}
1818

19-
cfg := &config.Config{
20-
Listeners: []config.ListenerConfig{{Name: "socks", Type: "socks5", Address: ":1080", Enabled: true}},
21-
}
22-
mgr = NewListenerManager(cfg)
23-
if _, ok := mgr.(NoopListenerManager); !ok {
24-
t.Fatalf("expected NoopListenerManager for non-http listeners, got %T", mgr)
25-
}
2619
}
2720

28-
func TestNewListenerManager_ConcreteForHTTP(t *testing.T) {
21+
func TestNewListenerManager_CompositeForEnabledListeners(t *testing.T) {
2922
t.Parallel()
3023

3124
cfg := &config.Config{
32-
Listeners: []config.ListenerConfig{{Name: "http", Type: "http", Address: "127.0.0.1:0", Enabled: true}},
25+
Listeners: []config.ListenerConfig{
26+
{Name: "http", Type: "http", Address: "127.0.0.1:0", Enabled: true},
27+
{Name: "socks", Type: "socks5", Address: "127.0.0.1:0", Enabled: true},
28+
},
3329
}
3430

3531
mgr := NewListenerManager(cfg)
36-
httpMgr, ok := mgr.(*HTTPListenerManager)
32+
composite, ok := mgr.(*CompositeListenerManager)
3733
if !ok {
38-
t.Fatalf("expected *HTTPListenerManager, got %T", mgr)
34+
t.Fatalf("expected *CompositeListenerManager, got %T", mgr)
35+
}
36+
if len(composite.managers) != 2 {
37+
t.Fatalf("expected 2 managers, got %d", len(composite.managers))
3938
}
4039

4140
ctx, cancel := context.WithTimeout(context.Background(), time.Second)
4241
defer cancel()
4342

44-
if err := httpMgr.Start(ctx); err != nil {
43+
if err := composite.Start(ctx); err != nil {
4544
t.Fatalf("start failed: %v", err)
4645
}
47-
if err := httpMgr.Shutdown(ctx); err != nil {
46+
if err := composite.Shutdown(ctx); err != nil {
4847
t.Fatalf("shutdown failed: %v", err)
4948
}
5049
}

0 commit comments

Comments
 (0)