chore(go.d/pkg/socket): add err to callback return values (#19103)
Ilya Mashchenko committed
Nov 29, 2024 at 12:43 UTC
61fb6d77d353e415b24bfc25f4190b1e23d7aa86
20 files changed
+477
-476
src/go/plugin/go.d/collector/beanstalk/client.go
+16
-24
@@ -88,9 +88,10 @@ func newBeanstalkConn(conf Config, log *logger.Logger) beanstalkConn {
88
return &beanstalkClient{
89
Logger: log,
90
client: socket.New(socket.Config{
91
- Address: conf.Address,
92
- Timeout: conf.Timeout.Duration(),
93
- TLSConf: nil,
91
+ Address: conf.Address,
92
+ Timeout: conf.Timeout.Duration(),
93
+ MaxReadLines: 2000,
94
+ TLSConf: nil,
95
}),
96
}
97
}
@@ -180,43 +181,34 @@ func (c *beanstalkClient) queryStatsTube(tubeName string) (*tubeStats, error) {
181
}
182
183
func (c *beanstalkClient) query(command string) (string, []byte, error) {
183
- var resp string
184
- var length int
185
- var body []byte
186
- var err error
187
-
184
c.Debugf("executing command: %s", command)
185
190
- const limitReadLines = 1000
191
- var num int
186
+ var (
187
+ resp string
188
+ body []byte
189
+ length int
190
+ err error
191
+ )
192
193
- clientErr := c.client.Command(command+"\r\n", func(line []byte) bool {
193
+ if err := c.client.Command(command+"\r\n", func(line []byte) (bool, error) {
194
if resp == "" {
195
s := string(line)
196
c.Debugf("command '%s' response: '%s'", command, s)
197
198
resp, length, err = parseResponseLine(s)
199
if err != nil {
200
- err = fmt.Errorf("command '%s' line '%s': %v", command, s, err)
200
+ return false, fmt.Errorf("command '%s' line '%s': %v", command, s, err)
201
}
202
- return err == nil && resp == "OK"
203
- }
202
205
- if num++; num >= limitReadLines {
206
- err = fmt.Errorf("command '%s': read line limit exceeded (%d)", command, limitReadLines)
207
- return false
203
+ return resp == "OK", nil
204
}
205
206
body = append(body, line...)
207
body = append(body, '\n')
208
213
- return len(body) < length
214
- })
215
- if clientErr != nil {
216
- return "", nil, fmt.Errorf("command '%s' client error: %v", command, clientErr)
217
- }
218
- if err != nil {
219
- return "", nil, err
209
+ return len(body) < length, nil
210
+ }); err != nil {
211
+ return "", nil, fmt.Errorf("command '%s': %v", command, err)
212
}
213
214
return resp, body, nil
src/go/plugin/go.d/collector/boinc/client.go
+5
-10
@@ -111,25 +111,20 @@ func (c *boincClient) send(req *boincRequest) (*boincReply, error) {
111
112
var b bytes.Buffer
113
114
- clientErr := c.conn.Command(string(reqData), func(bs []byte) bool {
114
+ if err := c.conn.Command(string(reqData), func(bs []byte) (bool, error) {
115
s := strings.TrimSpace(string(bs))
116
if s == "" {
117
- return true
117
+ return true, nil
118
}
119
120
if b.Len() == 0 && s != respStart {
121
- err = fmt.Errorf("unexpected response first line: %s", s)
122
- return false
121
+ return false, fmt.Errorf("unexpected response first line: %s", s)
122
}
123
124
b.WriteString(s)
125
127
- return s != respEnd
128
- })
129
- if clientErr != nil {
130
- return nil, fmt.Errorf("failed to send command: %v", clientErr)
131
- }
132
- if err != nil {
126
+ return s != respEnd, nil
127
+ }); err != nil {
128
return nil, fmt.Errorf("failed to send command: %v", err)
129
}
130
src/go/plugin/go.d/collector/dovecot/client.go
+3
-4
@@ -37,14 +37,13 @@ func (c *dovecotClient) queryExportGlobal() ([]byte, error) {
37
var b bytes.Buffer
38
var n int
39
40
- err := c.conn.Command("EXPORT\tglobal\n", func(bs []byte) bool {
40
+ if err := c.conn.Command("EXPORT\tglobal\n", func(bs []byte) (bool, error) {
41
b.Write(bs)
42
b.WriteByte('\n')
43
44
n++
45
- return n < 2
46
- })
47
- if err != nil {
45
+ return n < 2, nil
46
+ }); err != nil {
47
return nil, err
48
}
49
src/go/plugin/go.d/collector/gearman/client.go
+7
-18
@@ -19,8 +19,9 @@ type gearmanConn interface {
19
20
func newGearmanConn(conf Config) gearmanConn {
21
return &gearmanClient{conn: socket.New(socket.Config{
22
- Address: conf.Address,
23
- Timeout: conf.Timeout.Duration(),
22
+ Address: conf.Address,
23
+ Timeout: conf.Timeout.Duration(),
24
+ MaxReadLines: 10000,
25
})}
26
}
27
@@ -45,32 +46,20 @@ func (c *gearmanClient) queryPriorityStatus() ([]byte, error) {
46
}
47
48
func (c *gearmanClient) query(cmd string) ([]byte, error) {
48
- const limitReadLines = 10000
49
- var num int
50
- var err error
49
var b bytes.Buffer
50
53
- clientErr := c.conn.Command(cmd+"\n", func(bs []byte) bool {
51
+ if err := c.conn.Command(cmd+"\n", func(bs []byte) (bool, error) {
52
s := string(bs)
53
54
if strings.HasPrefix(s, "ERR") {
57
- err = fmt.Errorf("command '%s': %s", cmd, s)
58
- return false
55
+ return false, fmt.Errorf("command '%s': %s", cmd, s)
56
}
57
58
b.WriteString(s)
59
b.WriteByte('\n')
60
64
- if num++; num >= limitReadLines {
65
- err = fmt.Errorf("command '%s': read line limit exceeded (%d)", cmd, limitReadLines)
66
- return false
67
- }
68
- return !strings.HasPrefix(s, ".")
69
- })
70
- if clientErr != nil {
71
- return nil, fmt.Errorf("command '%s' client error: %v", cmd, clientErr)
72
- }
73
- if err != nil {
61
+ return !strings.HasPrefix(s, "."), nil
62
+ }); err != nil {
63
return nil, err
64
}
65
src/go/plugin/go.d/collector/hddtemp/client.go
+3
-9
@@ -25,21 +25,15 @@ type hddtempClient struct {
25
}
26
27
func (c *hddtempClient) queryHddTemp() (string, error) {
28
- var i int
29
- var s string
30
-
28
cfg := socket.Config{
29
Address: c.address,
30
Timeout: c.timeout,
31
}
32
36
- err := socket.ConnectAndRead(cfg, func(bs []byte) bool {
37
- if i++; i > 1 {
38
- return false
39
- }
33
+ var s string
34
+ err := socket.ConnectAndRead(cfg, func(bs []byte) (bool, error) {
35
s = string(bs)
41
- return true
42
-
36
+ return false, nil
37
})
38
if err != nil {
39
return "", err
src/go/plugin/go.d/collector/memcached/client.go
+4
-4
@@ -36,13 +36,13 @@ func (c *memcachedClient) disconnect() {
36
37
func (c *memcachedClient) queryStats() ([]byte, error) {
38
var b bytes.Buffer
39
- err := c.conn.Command("stats\r\n", func(bytes []byte) bool {
39
+ if err := c.conn.Command("stats\r\n", func(bytes []byte) (bool, error) {
40
s := strings.TrimSpace(string(bytes))
41
b.WriteString(s)
42
b.WriteByte('\n')
43
- return !(strings.HasPrefix(s, "END") || strings.HasPrefix(s, "ERROR"))
44
- })
45
- if err != nil {
43
+
44
+ return !(strings.HasPrefix(s, "END") || strings.HasPrefix(s, "ERROR")), nil
45
+ }); err != nil {
46
return nil, err
47
}
48
return b.Bytes(), nil
src/go/plugin/go.d/collector/openvpn/client/client.go
+7
-10
@@ -57,29 +57,26 @@ func (c *Client) Version() (*Version, error) {
57
58
func (c *Client) get(command string, stopRead stopReadFunc) (output []string, err error) {
59
var num int
60
- var maxLinesErr error
61
- err = c.Command(command, func(bytes []byte) bool {
60
+ if err := c.Command(command, func(bytes []byte) (bool, error) {
61
line := string(bytes)
62
num++
63
if num > maxLinesToRead {
65
- maxLinesErr = fmt.Errorf("read line limit exceeded (%d)", maxLinesToRead)
66
- return false
64
+ return false, fmt.Errorf("read line limit exceeded (%d)", maxLinesToRead)
65
}
66
67
// skip real-time messages
68
if strings.HasPrefix(line, ">") {
71
- return true
69
+ return true, nil
70
}
71
72
line = strings.Trim(line, "\r\n ")
73
output = append(output, line)
74
if stopRead != nil && stopRead(line) {
77
- return false
75
+ return false, nil
76
}
79
- return true
80
- })
81
- if maxLinesErr != nil {
82
- return nil, maxLinesErr
77
+ return true, nil
78
+ }); err != nil {
79
+ return nil, err
80
}
81
return output, err
82
}
src/go/plugin/go.d/collector/openvpn/client/client_test.go
+3
-1
@@ -98,7 +98,9 @@ func (m *mockSocketClient) Command(command string, process socket.Processor) err
98
}
99
100
for s.Scan() {
101
- process(s.Bytes())
101
+ if _, err := process(s.Bytes()); err != nil {
102
+ return err
103
+ }
104
}
105
return nil
106
}
src/go/plugin/go.d/collector/tor/client.go
+8
-14
@@ -58,9 +58,9 @@ func (c *torControlClient) authenticate() error {
58
}
59
60
var s string
61
- err := c.conn.Command(cmd+"\n", func(bs []byte) bool {
61
+ err := c.conn.Command(cmd+"\n", func(bs []byte) (bool, error) {
62
s = string(bs)
63
- return false
63
+ return false, nil
64
})
65
if err != nil {
66
return fmt.Errorf("authentication failed: %v", err)
@@ -74,7 +74,7 @@ func (c *torControlClient) authenticate() error {
74
func (c *torControlClient) disconnect() {
75
// https://spec.torproject.org/control-spec/commands.html#quit
76
77
- _ = c.conn.Command(cmdQuit+"\n", func(bs []byte) bool { return false })
77
+ _ = c.conn.Command(cmdQuit+"\n", func(bs []byte) (bool, error) { return false, nil })
78
_ = c.conn.Disconnect()
79
}
80
@@ -87,27 +87,21 @@ func (c *torControlClient) getInfo(keywords ...string) ([]byte, error) {
87
cmd := fmt.Sprintf("%s %s", cmdGetInfo, strings.Join(keywords, " "))
88
89
var buf bytes.Buffer
90
- var err error
90
92
- clientErr := c.conn.Command(cmd+"\n", func(bs []byte) bool {
91
+ if err := c.conn.Command(cmd+"\n", func(bs []byte) (bool, error) {
92
s := string(bs)
93
94
switch {
95
case strings.HasPrefix(s, "250-"):
96
buf.WriteString(strings.TrimPrefix(s, "250-"))
97
buf.WriteByte('\n')
99
- return true
98
+ return true, nil
99
case strings.HasPrefix(s, "250 "):
101
- return false
100
+ return false, nil
101
default:
103
- err = errors.New(s)
104
- return false
102
+ return false, errors.New(s)
103
}
106
- })
107
- if clientErr != nil {
108
- return nil, fmt.Errorf("command '%s' failed: %v", cmd, clientErr)
109
- }
110
- if err != nil {
104
+ }); err != nil {
105
return nil, fmt.Errorf("command '%s' failed: %v", cmd, err)
106
}
107
src/go/plugin/go.d/collector/unbound/collect.go
+2
-2
@@ -36,9 +36,9 @@ func (c *Collector) scrapeUnboundStats() ([]entry, error) {
36
}
37
defer func() { _ = c.client.Disconnect() }()
38
39
- err := c.client.Command(command+"\n", func(bytes []byte) bool {
39
+ err := c.client.Command(command+"\n", func(bytes []byte) (bool, error) {
40
output = append(output, string(bytes))
41
- return true
41
+ return true, nil
42
})
43
if err != nil {
44
return nil, fmt.Errorf("send command '%s': %w", command, err)
src/go/plugin/go.d/collector/upsd/client.go
+2
-2
@@ -133,7 +133,7 @@ func (c *upsdClient) sendCommand(cmd string) ([]string, error) {
133
var errMsg string
134
endLine := getEndLine(cmd)
135
136
- err := c.conn.Command(cmd+"\n", func(bytes []byte) bool {
136
+ err := c.conn.Command(cmd+"\n", func(bytes []byte) (bool, error) {
137
line := string(bytes)
138
resp = append(resp, line)
139
@@ -141,7 +141,7 @@ func (c *upsdClient) sendCommand(cmd string) ([]string, error) {
141
errMsg = strings.TrimPrefix(line, "ERR ")
142
}
143
144
- return line != endLine && errMsg == ""
144
+ return line != endLine && errMsg == "", nil
145
})
146
if err != nil {
147
return nil, err
src/go/plugin/go.d/collector/uwsgi/client.go
+6
-18
@@ -4,7 +4,6 @@ package uwsgi
4
5
import (
6
"bytes"
7
- "fmt"
7
"time"
8
9
"github.com/netdata/netdata/go/plugins/plugin/go.d/pkg/socket"
@@ -28,30 +27,19 @@ type uwsgiClient struct {
27
28
func (c *uwsgiClient) queryStats() ([]byte, error) {
29
var b bytes.Buffer
31
- var n int64
32
- var err error
33
- const readLineLimit = 1000 * 10
30
31
cfg := socket.Config{
36
- Address: c.address,
37
- Timeout: c.timeout,
32
+ Address: c.address,
33
+ Timeout: c.timeout,
34
+ MaxReadLines: 1000 * 10,
35
}
36
40
- clientErr := socket.ConnectAndRead(cfg, func(bs []byte) bool {
37
+ if err := socket.ConnectAndRead(cfg, func(bs []byte) (bool, error) {
38
b.Write(bs)
39
b.WriteByte('\n')
43
-
44
- if n++; n >= readLineLimit {
45
- err = fmt.Errorf("read line limit exceeded %d", readLineLimit)
46
- return false
47
- }
40
// The server will close the connection when it has finished sending data.
49
- return true
50
- })
51
- if clientErr != nil {
52
- return nil, clientErr
53
- }
54
- if err != nil {
41
+ return true, nil
42
+ }); err != nil {
43
return nil, err
44
}
45
src/go/plugin/go.d/collector/zookeeper/fetcher.go
+3
-15
@@ -4,14 +4,11 @@ package zookeeper
4
5
import (
6
"bytes"
7
- "fmt"
7
"unsafe"
8
9
"github.com/netdata/netdata/go/plugins/plugin/go.d/pkg/socket"
10
)
11
13
-const limitReadLines = 2000
14
-
12
type fetcher interface {
13
fetch(command string) ([]string, error)
14
}
@@ -26,21 +23,12 @@ func (c *zookeeperFetcher) fetch(command string) (rows []string, err error) {
23
}
24
defer func() { _ = c.Disconnect() }()
25
29
- var num int
30
- clientErr := c.Command(command, func(b []byte) bool {
26
+ if err := c.Command(command, func(b []byte) (bool, error) {
27
if !isZKLine(b) || isMntrLineOK(b) {
28
rows = append(rows, string(b))
29
}
34
- if num += 1; num >= limitReadLines {
35
- err = fmt.Errorf("read line limit exceeded (%d)", limitReadLines)
36
- return false
37
- }
38
- return true
39
- })
40
- if clientErr != nil {
41
- return nil, clientErr
42
- }
43
- if err != nil {
30
+ return true, nil
31
+ }); err != nil {
32
return nil, err
33
}
34
src/go/plugin/go.d/collector/zookeeper/fetcher_test.go
+3
-9
@@ -22,14 +22,6 @@ func Test_clientFetch(t *testing.T) {
22
assert.Len(t, rows, 10)
23
}
24
25
-func Test_clientFetchReadLineLimitExceeded(t *testing.T) {
26
- c := &zookeeperFetcher{Client: &mockSocket{rowsNumResp: limitReadLines + 1}}
27
-
28
- rows, err := c.fetch("whatever\n")
29
- assert.Error(t, err)
30
- assert.Len(t, rows, 0)
31
-}
32
-
25
type mockSocket struct {
26
rowsNumResp int
27
}
@@ -44,7 +36,9 @@ func (m *mockSocket) Disconnect() error {
36
37
func (m *mockSocket) Command(command string, process socket.Processor) error {
38
for i := 0; i < m.rowsNumResp; i++ {
47
- process([]byte(command))
39
+ if _, err := process([]byte(command)); err != nil {
40
+ return err
41
+ }
42
}
43
return nil
44
}
src/go/plugin/go.d/collector/zookeeper/init.go
+4
-3
@@ -30,9 +30,10 @@ func (c *Collector) initZookeeperFetcher() (fetcher, error) {
30
}
31
32
sock := socket.New(socket.Config{
33
- Address: c.Address,
34
- Timeout: c.Timeout.Duration(),
35
- TLSConf: tlsConf,
33
+ Address: c.Address,
34
+ Timeout: c.Timeout.Duration(),
35
+ TLSConf: tlsConf,
36
+ MaxReadLines: 2000,
37
})
38
39
return &zookeeperFetcher{Client: sock}, nil
src/go/plugin/go.d/pkg/socket/client.go
+67
-51
@@ -4,25 +4,29 @@ package socket
4
5
import (
6
"bufio"
7
+ "context"
8
"crypto/tls"
9
"errors"
10
+ "fmt"
11
"net"
12
"time"
13
)
14
13
-// Processor function passed to the Socket.Command function.
14
-// It is passed by the caller to process a command's response line by line.
15
-type Processor func([]byte) bool
15
+// Processor is a callback function passed to the Socket.Command method.
16
+// It processes each response line received from the server.
17
+type Processor func([]byte) (bool, error)
18
17
-// Client is the interface that wraps the basic socket client operations
18
-// and hides the implementation details from the users.
19
-// Implementations should return TCP, UDP or Unix ready sockets.
19
+// Client defines an interface for socket clients, abstracting the underlying implementation.
20
+// Implementations should provide connections for various socket types such as TCP, UDP, or Unix domain sockets.
21
type Client interface {
22
Connect() error
23
Disconnect() error
24
Command(command string, process Processor) error
25
}
26
27
+// ConnectAndRead establishes a connection using the given configuration,
28
+// executes the provided processor function on the incoming response lines,
29
+// and ensures the connection is properly closed after use.
30
func ConnectAndRead(cfg Config, process Processor) error {
31
sock := New(cfg)
32
@@ -35,46 +39,33 @@ func ConnectAndRead(cfg Config, process Processor) error {
39
return sock.read(process)
40
}
41
38
-// New returns a new pointer to a socket client given the socket
39
-// type (IP, TCP, UDP, UNIX), a network address (IP/domain:port),
40
-// a timeout and a TLS config. It supports both IPv4 and IPv6 address
41
-// and reuses connection where possible.
42
+// New creates and returns a new Socket instance configured with the provided settings.
43
+// The socket supports multiple types (TCP, UDP, UNIX), addresses (IPv4, IPv6, domain names),
44
+// and optional TLS encryption. Connections are reused where possible.
45
func New(cfg Config) *Socket {
46
return &Socket{Config: cfg}
47
}
48
46
-// Socket is the implementation of a socket client.
49
+// Socket is a concrete implementation of the Client interface, managing a network connection
50
+// based on the specified configuration (address, type, timeout, and optional TLS settings).
51
type Socket struct {
52
Config
53
conn net.Conn
54
}
55
52
-// Config holds the network ip v4 or v6 address, port,
53
-// Socket type(ip, tcp, udp, unix), timeout and TLS configuration for a Socket
56
+// Config encapsulates the settings required to establish a network connection.
57
type Config struct {
55
- Address string
56
- Timeout time.Duration
57
- TLSConf *tls.Config
58
+ Address string
59
+ Timeout time.Duration
60
+ TLSConf *tls.Config
61
+ MaxReadLines int64
62
}
63
60
-// Connect connects to the Socket address on the named network.
61
-// If the address is a domain name it will also perform the DNS resolution.
62
-// Address like :80 will attempt to connect to the localhost.
63
-// The config timeout and TLS config will be used.
64
+// Connect establishes a connection to the specified address using the configuration details.
65
func (s *Socket) Connect() error {
65
- network, address := networkType(s.Address)
66
- var conn net.Conn
67
- var err error
68
-
69
- if s.TLSConf == nil {
70
- conn, err = net.DialTimeout(network, address, s.timeout())
71
- } else {
72
- var d net.Dialer
73
- d.Timeout = s.timeout()
74
- conn, err = tls.DialWithDialer(&d, network, address, s.TLSConf)
75
- }
66
+ conn, err := s.dial()
67
if err != nil {
77
- return err
68
+ return fmt.Errorf("socket.Connect: %w", err)
69
}
70
71
s.conn = conn
@@ -82,22 +73,19 @@ func (s *Socket) Connect() error {
73
return nil
74
}
75
85
-// Disconnect closes the connection.
86
-// Any in-flight commands will be cancelled and return errors.
87
-func (s *Socket) Disconnect() (err error) {
88
- if s.conn != nil {
89
- err = s.conn.Close()
90
- s.conn = nil
76
+// Disconnect terminates the active connection if one exists.
77
+func (s *Socket) Disconnect() error {
78
+ if s.conn == nil {
79
+ return nil
80
}
81
+ err := s.conn.Close()
82
+ s.conn = nil
83
return err
84
}
85
95
-// Command writes the command string to the connection and passed the
96
-// response bytes line by line to the process function. It uses the
97
-// timeout value from the Socket config and returns read, write and
98
-// timeout errors if any. If a timeout occurs during the processing
99
-// of the responses this function will stop processing and return a
100
-// timeout error.
86
+// Command sends a command string to the connected server and processes its response line by line
87
+// using the provided Processor function. This method respects the timeout configuration
88
+// for write and read operations. If a timeout or processing error occurs, it stops and returns the error.
89
func (s *Socket) Command(command string, process Processor) error {
90
if s.conn == nil {
91
return errors.New("cannot send command on nil connection")
@@ -112,10 +100,10 @@ func (s *Socket) Command(command string, process Processor) error {
100
101
func (s *Socket) write(command string) error {
102
if s.conn == nil {
115
- return errors.New("attempt to write on nil connection")
103
+ return errors.New("write: nil connection")
104
}
105
118
- if err := s.conn.SetWriteDeadline(time.Now().Add(s.timeout())); err != nil {
106
+ if err := s.conn.SetWriteDeadline(s.deadline()); err != nil {
107
return err
108
}
109
@@ -126,25 +114,53 @@ func (s *Socket) write(command string) error {
114
115
func (s *Socket) read(process Processor) error {
116
if process == nil {
129
- return errors.New("process func is nil")
117
+ return errors.New("read: process func is nil")
118
}
131
-
119
if s.conn == nil {
133
- return errors.New("attempt to read on nil connection")
120
+ return errors.New("read: nil connection")
121
}
122
136
- if err := s.conn.SetReadDeadline(time.Now().Add(s.timeout())); err != nil {
123
+ if err := s.conn.SetReadDeadline(s.deadline()); err != nil {
124
return err
125
}
126
127
sc := bufio.NewScanner(s.conn)
128
142
- for sc.Scan() && process(sc.Bytes()) {
129
+ var n int64
130
+ limit := s.MaxReadLines
131
+
132
+ for sc.Scan() {
133
+ more, err := process(sc.Bytes())
134
+ if err != nil {
135
+ return err
136
+ }
137
+ if n++; limit > 0 && n > limit {
138
+ return fmt.Errorf("read line limit exceeded (%d", limit)
139
+ }
140
+ if !more {
141
+ break
142
+ }
143
}
144
145
return sc.Err()
146
}
147
148
+func (s *Socket) dial() (net.Conn, error) {
149
+ network, address := parseAddress(s.Address)
150
+
151
+ var d net.Dialer
152
+ d.Timeout = s.timeout()
153
+
154
+ if s.TLSConf != nil {
155
+ return tls.DialWithDialer(&d, network, address, s.TLSConf)
156
+ }
157
+ return d.DialContext(context.Background(), network, address)
158
+}
159
+
160
+func (s *Socket) deadline() time.Time {
161
+ return time.Now().Add(s.timeout())
162
+}
163
+
164
func (s *Socket) timeout() time.Duration {
165
if s.Timeout == 0 {
166
return time.Second
src/go/plugin/go.d/pkg/socket/client_test.go
+76
-142
@@ -3,152 +3,86 @@
3
package socket
4
5
import (
6
- "crypto/tls"
6
"testing"
7
"time"
8
10
- "github.com/stretchr/testify/assert"
9
"github.com/stretchr/testify/require"
10
)
11
14
-const (
15
- testServerAddress = "127.0.0.1:9999"
16
- testUdpServerAddress = "udp://127.0.0.1:9999"
17
- testUnixServerAddress = "/tmp/testSocketFD"
18
- defaultTimeout = 100 * time.Millisecond
19
-)
20
-
21
-var tcpConfig = Config{
22
- Address: testServerAddress,
23
- Timeout: defaultTimeout,
24
- TLSConf: nil,
25
-}
26
-
27
-var udpConfig = Config{
28
- Address: testUdpServerAddress,
29
- Timeout: defaultTimeout,
30
- TLSConf: nil,
31
-}
32
-
33
-var unixConfig = Config{
34
- Address: testUnixServerAddress,
35
- Timeout: defaultTimeout,
36
- TLSConf: nil,
37
-}
38
-
39
-var tcpTlsConfig = Config{
40
- Address: testServerAddress,
41
- Timeout: defaultTimeout,
42
- TLSConf: &tls.Config{},
43
-}
44
-
45
-func Test_clientCommand(t *testing.T) {
46
- srv := &tcpServer{addr: testServerAddress, rowsNumResp: 1}
47
- go func() { _ = srv.Run(); defer func() { _ = srv.Close() }() }()
48
-
49
- time.Sleep(time.Millisecond * 100)
50
- sock := New(tcpConfig)
51
- require.NoError(t, sock.Connect())
52
- err := sock.Command("ping\n", func(bytes []byte) bool {
53
- assert.Equal(t, "pong", string(bytes))
54
- return true
55
- })
56
- require.NoError(t, sock.Disconnect())
57
- require.NoError(t, err)
58
-}
59
-
60
-func Test_clientTimeout(t *testing.T) {
61
- srv := &tcpServer{addr: testServerAddress, rowsNumResp: 1}
62
- go func() { _ = srv.Run() }()
63
-
64
- time.Sleep(time.Millisecond * 100)
65
- sock := New(tcpConfig)
66
- require.NoError(t, sock.Connect())
67
- sock.Timeout = 0
68
- err := sock.Command("ping\n", func(bytes []byte) bool {
69
- assert.Equal(t, "pong", string(bytes))
70
- return true
71
- })
72
- require.NoError(t, err)
73
-}
74
-
75
-func Test_clientIncompleteSSL(t *testing.T) {
76
- srv := &tcpServer{addr: testServerAddress, rowsNumResp: 1}
77
- go func() { _ = srv.Run() }()
78
-
79
- time.Sleep(time.Millisecond * 100)
80
- sock := New(tcpTlsConfig)
81
- err := sock.Connect()
82
- require.Error(t, err)
83
-}
84
-
85
-func Test_clientCommandStopProcessing(t *testing.T) {
86
- srv := &tcpServer{addr: testServerAddress, rowsNumResp: 2}
87
- go func() { _ = srv.Run() }()
88
-
89
- time.Sleep(time.Millisecond * 100)
90
- sock := New(tcpConfig)
91
- require.NoError(t, sock.Connect())
92
- err := sock.Command("ping\n", func(bytes []byte) bool {
93
- assert.Equal(t, "pong", string(bytes))
94
- return false
95
- })
96
- require.NoError(t, sock.Disconnect())
97
- require.NoError(t, err)
98
-}
99
-
100
-func Test_clientUDPCommand(t *testing.T) {
101
- srv := &udpServer{addr: testServerAddress, rowsNumResp: 1}
102
- go func() { _ = srv.Run(); defer func() { _ = srv.Close() }() }()
103
-
104
- time.Sleep(time.Millisecond * 100)
105
- sock := New(udpConfig)
106
- require.NoError(t, sock.Connect())
107
- err := sock.Command("ping\n", func(bytes []byte) bool {
108
- assert.Equal(t, "pong", string(bytes))
109
- return false
110
- })
111
- require.NoError(t, sock.Disconnect())
112
- require.NoError(t, err)
113
-}
114
-
115
-func Test_clientTCPAddress(t *testing.T) {
116
- srv := &tcpServer{addr: testServerAddress, rowsNumResp: 1}
117
- go func() { _ = srv.Run() }()
118
- time.Sleep(time.Millisecond * 100)
119
-
120
- sock := New(tcpConfig)
121
- require.NoError(t, sock.Connect())
122
-
123
- tcpConfig.Address = "tcp://" + tcpConfig.Address
124
- sock = New(tcpConfig)
125
- require.NoError(t, sock.Connect())
126
-}
127
-
128
-func Test_clientUnixCommand(t *testing.T) {
129
- srv := &unixServer{addr: testUnixServerAddress, rowsNumResp: 1}
130
- // cleanup previous file descriptors
131
- _ = srv.Close()
132
- go func() { _ = srv.Run() }()
133
-
134
- time.Sleep(time.Millisecond * 200)
135
- sock := New(unixConfig)
136
- require.NoError(t, sock.Connect())
137
- err := sock.Command("ping\n", func(bytes []byte) bool {
138
- assert.Equal(t, "pong", string(bytes))
139
- return false
140
- })
141
- require.NoError(t, err)
142
- require.NoError(t, sock.Disconnect())
143
-}
144
-
145
-func Test_clientEmptyProcessFunc(t *testing.T) {
146
- srv := &tcpServer{addr: testServerAddress, rowsNumResp: 1}
147
- go func() { _ = srv.Run() }()
148
-
149
- time.Sleep(time.Millisecond * 100)
150
- sock := New(tcpConfig)
151
- require.NoError(t, sock.Connect())
152
- err := sock.Command("ping\n", nil)
153
- require.Error(t, err, "nil process func should return an error")
12
+func TestSocket_Command(t *testing.T) {
13
+ const (
14
+ testServerAddress = "tcp://127.0.0.1:9999"
15
+ testUdpServerAddress = "udp://127.0.0.1:9999"
16
+ testUnixServerAddress = "unix:///tmp/testSocketFD"
17
+ defaultTimeout = 1000 * time.Millisecond
18
+ )
19
+
20
+ type server interface {
21
+ Run() error
22
+ Close() error
23
+ }
24
+
25
+ tests := map[string]struct {
26
+ srv server
27
+ cfg Config
28
+ wantConnectErr bool
29
+ wantCommandErr bool
30
+ }{
31
+ "tcp": {
32
+ srv: newTCPServer(testServerAddress),
33
+ cfg: Config{
34
+ Address: testServerAddress,
35
+ Timeout: defaultTimeout,
36
+ },
37
+ },
38
+ "udp": {
39
+ srv: newUDPServer(testUdpServerAddress),
40
+ cfg: Config{
41
+ Address: testUdpServerAddress,
42
+ Timeout: defaultTimeout,
43
+ },
44
+ },
45
+ "unix": {
46
+ srv: newUnixServer(testUnixServerAddress),
47
+ cfg: Config{
48
+ Address: testUnixServerAddress,
49
+ Timeout: defaultTimeout,
50
+ },
51
+ },
52
+ }
53
+
54
+ for name, test := range tests {
55
+ t.Run(name, func(t *testing.T) {
56
+ go func() {
57
+ defer func() { _ = test.srv.Close() }()
58
+ require.NoError(t, test.srv.Run())
59
+ }()
60
+ time.Sleep(time.Millisecond * 500)
61
+
62
+ sock := New(test.cfg)
63
+
64
+ err := sock.Connect()
65
+
66
+ if test.wantConnectErr {
67
+ require.Error(t, err)
68
+ return
69
+ }
70
+ require.NoError(t, err)
71
+
72
+ defer sock.Disconnect()
73
+
74
+ var resp string
75
+ err = sock.Command("ping\n", func(bytes []byte) (bool, error) {
76
+ resp = string(bytes)
77
+ return false, nil
78
+ })
79
+
80
+ if test.wantCommandErr {
81
+ require.Error(t, err)
82
+ } else {
83
+ require.NoError(t, err)
84
+ require.Equal(t, "pong", resp)
85
+ }
86
+ })
87
+ }
88
}
src/go/plugin/go.d/pkg/socket/server.go
new
+257
@@ -0,0 +1,257 @@
1
+// SPDX-License-Identifier: GPL-3.0-or-later
2
+
3
+package socket
4
+
5
+import (
6
+ "bufio"
7
+ "context"
8
+ "errors"
9
+ "fmt"
10
+ "net"
11
+ "os"
12
+ "sync"
13
+ "time"
14
+)
15
+
16
+func newTCPServer(addr string) *tcpServer {
17
+ ctx, cancel := context.WithCancel(context.Background())
18
+ _, addr = parseAddress(addr)
19
+ return &tcpServer{
20
+ addr: addr,
21
+ ctx: ctx,
22
+ cancel: cancel,
23
+ }
24
+}
25
+
26
+type tcpServer struct {
27
+ addr string
28
+ listener net.Listener
29
+ wg sync.WaitGroup
30
+ ctx context.Context
31
+ cancel context.CancelFunc
32
+}
33
+
34
+func (t *tcpServer) Run() error {
35
+ var err error
36
+ t.listener, err = net.Listen("tcp", t.addr)
37
+ if err != nil {
38
+ return fmt.Errorf("failed to start TCP server: %w", err)
39
+ }
40
+ return t.handleConnections()
41
+}
42
+
43
+func (t *tcpServer) Close() (err error) {
44
+ t.cancel()
45
+ if t.listener != nil {
46
+ if err := t.listener.Close(); err != nil {
47
+ return fmt.Errorf("failed to close TCP server: %w", err)
48
+ }
49
+ }
50
+ t.wg.Wait()
51
+ return nil
52
+}
53
+
54
+func (t *tcpServer) handleConnections() (err error) {
55
+ for {
56
+ select {
57
+ case <-t.ctx.Done():
58
+ return nil
59
+ default:
60
+ conn, err := t.listener.Accept()
61
+ if err != nil {
62
+ if errors.Is(err, net.ErrClosed) {
63
+ return nil
64
+ }
65
+ return fmt.Errorf("could not accept connection: %v", err)
66
+ }
67
+ t.wg.Add(1)
68
+ go func() {
69
+ defer t.wg.Done()
70
+ t.handleConnection(conn)
71
+ }()
72
+ }
73
+ }
74
+}
75
+
76
+func (t *tcpServer) handleConnection(conn net.Conn) {
77
+ defer func() { _ = conn.Close() }()
78
+
79
+ if err := conn.SetDeadline(time.Now().Add(time.Second)); err != nil {
80
+ return
81
+ }
82
+
83
+ rw := bufio.NewReadWriter(bufio.NewReader(conn), bufio.NewWriter(conn))
84
+
85
+ if _, err := rw.ReadString('\n'); err != nil {
86
+ writeResponse(rw, fmt.Sprintf("failed to read input: %v\n", err))
87
+ } else {
88
+ writeResponse(rw, "pong\n")
89
+ }
90
+}
91
+
92
+func newUDPServer(addr string) *udpServer {
93
+ ctx, cancel := context.WithCancel(context.Background())
94
+ _, addr = parseAddress(addr)
95
+ return &udpServer{
96
+ addr: addr,
97
+ ctx: ctx,
98
+ cancel: cancel,
99
+ }
100
+}
101
+
102
+type udpServer struct {
103
+ addr string
104
+ conn *net.UDPConn
105
+ ctx context.Context
106
+ cancel context.CancelFunc
107
+}
108
+
109
+func (u *udpServer) Run() error {
110
+ addr, err := net.ResolveUDPAddr("udp", u.addr)
111
+ if err != nil {
112
+ return fmt.Errorf("failed to resolve UDP address: %w", err)
113
+ }
114
+
115
+ u.conn, err = net.ListenUDP("udp", addr)
116
+ if err != nil {
117
+ return fmt.Errorf("failed to start UDP server: %w", err)
118
+ }
119
+
120
+ return u.handleConnections()
121
+}
122
+
123
+func (u *udpServer) Close() (err error) {
124
+ u.cancel()
125
+ if u.conn != nil {
126
+ if err := u.conn.Close(); err != nil {
127
+ return fmt.Errorf("failed to close UDP server: %w", err)
128
+ }
129
+ }
130
+ return nil
131
+}
132
+
133
+func (u *udpServer) handleConnections() error {
134
+ buffer := make([]byte, 8192)
135
+ for {
136
+ select {
137
+ case <-u.ctx.Done():
138
+ return nil
139
+ default:
140
+ if err := u.conn.SetReadDeadline(time.Now().Add(time.Second)); err != nil {
141
+ continue
142
+ }
143
+
144
+ _, addr, err := u.conn.ReadFromUDP(buffer[0:])
145
+ if err != nil {
146
+ if !errors.Is(err, os.ErrDeadlineExceeded) {
147
+ return fmt.Errorf("failed to read UDP packet: %w", err)
148
+ }
149
+ continue
150
+ }
151
+
152
+ if _, err := u.conn.WriteToUDP([]byte("pong\n"), addr); err != nil {
153
+ return fmt.Errorf("failed to write UDP response: %w", err)
154
+ }
155
+ }
156
+ }
157
+}
158
+
159
+func newUnixServer(addr string) *unixServer {
160
+ ctx, cancel := context.WithCancel(context.Background())
161
+ _, addr = parseAddress(addr)
162
+ return &unixServer{
163
+ addr: addr,
164
+ ctx: ctx,
165
+ cancel: cancel,
166
+ }
167
+}
168
+
169
+type unixServer struct {
170
+ addr string
171
+ listener *net.UnixListener
172
+ wg sync.WaitGroup
173
+ ctx context.Context
174
+ cancel context.CancelFunc
175
+}
176
+
177
+func (u *unixServer) Run() error {
178
+ if err := os.Remove(u.addr); err != nil && !os.IsNotExist(err) {
179
+ return fmt.Errorf("failed to clean up existing socket: %w", err)
180
+ }
181
+
182
+ addr, err := net.ResolveUnixAddr("unix", u.addr)
183
+ if err != nil {
184
+ return fmt.Errorf("failed to resolve Unix address: %w", err)
185
+ }
186
+
187
+ u.listener, err = net.ListenUnix("unix", addr)
188
+ if err != nil {
189
+ return fmt.Errorf("failed to start Unix server: %w", err)
190
+ }
191
+
192
+ return u.handleConnections()
193
+}
194
+
195
+func (u *unixServer) Close() error {
196
+ u.cancel()
197
+
198
+ if u.listener != nil {
199
+ if err := u.listener.Close(); err != nil {
200
+ return fmt.Errorf("failed to close Unix server: %w", err)
201
+ }
202
+ }
203
+
204
+ u.wg.Wait()
205
+ _ = os.Remove(u.addr)
206
+
207
+ return nil
208
+}
209
+
210
+func (u *unixServer) handleConnections() error {
211
+ for {
212
+ select {
213
+ case <-u.ctx.Done():
214
+ return nil
215
+ default:
216
+ if err := u.listener.SetDeadline(time.Now().Add(time.Second)); err != nil {
217
+ continue
218
+ }
219
+
220
+ conn, err := u.listener.AcceptUnix()
221
+ if err != nil {
222
+ if !errors.Is(err, os.ErrDeadlineExceeded) {
223
+ return err
224
+ }
225
+ continue
226
+ }
227
+
228
+ u.wg.Add(1)
229
+ go func() {
230
+ defer u.wg.Done()
231
+ u.handleConnection(conn)
232
+ }()
233
+ }
234
+ }
235
+}
236
+
237
+func (u *unixServer) handleConnection(conn net.Conn) {
238
+ defer func() { _ = conn.Close() }()
239
+
240
+ if err := conn.SetDeadline(time.Now().Add(time.Second)); err != nil {
241
+ return
242
+ }
243
+
244
+ rw := bufio.NewReadWriter(bufio.NewReader(conn), bufio.NewWriter(conn))
245
+
246
+ if _, err := rw.ReadString('\n'); err != nil {
247
+ writeResponse(rw, fmt.Sprintf("failed to read input: %v\n", err))
248
+ } else {
249
+ writeResponse(rw, "pong\n")
250
+ }
251
+
252
+}
253
+
254
+func writeResponse(rw *bufio.ReadWriter, response string) {
255
+ _, _ = rw.WriteString(response)
256
+ _ = rw.Flush()
257
+}
src/go/plugin/go.d/pkg/socket/servers_test.go
deleted
-139
@@ -1,139 +0,0 @@
1
-// SPDX-License-Identifier: GPL-3.0-or-later
2
-
3
-package socket
4
-
5
-import (
6
- "bufio"
7
- "errors"
8
- "fmt"
9
- "net"
10
- "os"
11
- "strings"
12
- "time"
13
-)
14
-
15
-type tcpServer struct {
16
- addr string
17
- server net.Listener
18
- rowsNumResp int
19
-}
20
-
21
-func (t *tcpServer) Run() (err error) {
22
- t.server, err = net.Listen("tcp", t.addr)
23
- if err != nil {
24
- return
25
- }
26
- return t.handleConnections()
27
-}
28
-
29
-func (t *tcpServer) Close() (err error) {
30
- return t.server.Close()
31
-}
32
-
33
-func (t *tcpServer) handleConnections() (err error) {
34
- for {
35
- conn, err := t.server.Accept()
36
- if err != nil || conn == nil {
37
- return errors.New("could not accept connection")
38
- }
39
- t.handleConnection(conn)
40
- }
41
-}
42
-
43
-func (t *tcpServer) handleConnection(conn net.Conn) {
44
- defer func() { _ = conn.Close() }()
45
- _ = conn.SetDeadline(time.Now().Add(time.Millisecond * 100))
46
-
47
- rw := bufio.NewReadWriter(bufio.NewReader(conn), bufio.NewWriter(conn))
48
- _, err := rw.ReadString('\n')
49
- if err != nil {
50
- _, _ = rw.WriteString("failed to read input")
51
- _ = rw.Flush()
52
- } else {
53
- resp := strings.Repeat("pong\n", t.rowsNumResp)
54
- _, _ = rw.WriteString(resp)
55
- _ = rw.Flush()
56
- }
57
-}
58
-
59
-type udpServer struct {
60
- addr string
61
- conn *net.UDPConn
62
- rowsNumResp int
63
-}
64
-
65
-func (u *udpServer) Run() (err error) {
66
- addr, err := net.ResolveUDPAddr("udp", u.addr)
67
- if err != nil {
68
- return err
69
- }
70
- u.conn, err = net.ListenUDP("udp", addr)
71
- if err != nil {
72
- return
73
- }
74
- u.handleConnections()
75
- return nil
76
-}
77
-
78
-func (u *udpServer) Close() (err error) {
79
- return u.conn.Close()
80
-}
81
-
82
-func (u *udpServer) handleConnections() {
83
- for {
84
- var buf [2048]byte
85
- _, addr, _ := u.conn.ReadFromUDP(buf[0:])
86
- resp := strings.Repeat("pong\n", u.rowsNumResp)
87
- _, _ = u.conn.WriteToUDP([]byte(resp), addr)
88
- }
89
-}
90
-
91
-type unixServer struct {
92
- addr string
93
- conn *net.UnixListener
94
- rowsNumResp int
95
-}
96
-
97
-func (u *unixServer) Run() (err error) {
98
- _, _ = os.CreateTemp("/tmp", "testSocketFD")
99
- addr, err := net.ResolveUnixAddr("unix", u.addr)
100
- if err != nil {
101
- return err
102
- }
103
- u.conn, err = net.ListenUnix("unix", addr)
104
- if err != nil {
105
- return
106
- }
107
- go u.handleConnections()
108
- return nil
109
-}
110
-
111
-func (u *unixServer) Close() (err error) {
112
- _ = os.Remove(testUnixServerAddress)
113
- return u.conn.Close()
114
-}
115
-
116
-func (u *unixServer) handleConnections() {
117
- var conn net.Conn
118
- var err error
119
- conn, err = u.conn.AcceptUnix()
120
- if err != nil {
121
- panic(fmt.Errorf("could not accept connection: %v", err))
122
- }
123
- u.handleConnection(conn)
124
-}
125
-
126
-func (u *unixServer) handleConnection(conn net.Conn) {
127
- _ = conn.SetDeadline(time.Now().Add(time.Second))
128
-
129
- rw := bufio.NewReadWriter(bufio.NewReader(conn), bufio.NewWriter(conn))
130
- _, err := rw.ReadString('\n')
131
- if err != nil {
132
- _, _ = rw.WriteString("failed to read input")
133
- _ = rw.Flush()
134
- } else {
135
- resp := strings.Repeat("pong\n", u.rowsNumResp)
136
- _, _ = rw.WriteString(resp)
137
- _ = rw.Flush()
138
- }
139
-}
src/go/plugin/go.d/pkg/socket/utils.go
+1
-1
@@ -12,7 +12,7 @@ func IsUdpSocket(address string) bool {
12
return strings.HasPrefix(address, "udp://")
13
}
14
15
-func networkType(address string) (string, string) {
15
+func parseAddress(address string) (string, string) {
16
switch {
17
case IsUnixSocket(address):
18
address = strings.TrimPrefix(address, "unix://")