-
Notifications
You must be signed in to change notification settings - Fork 0
Expand file tree
/
Copy pathipc.go
More file actions
258 lines (227 loc) · 5.82 KB
/
Copy pathipc.go
File metadata and controls
258 lines (227 loc) · 5.82 KB
1
2
3
4
5
6
7
8
9
10
11
12
13
14
15
16
17
18
19
20
21
22
23
24
25
26
27
28
29
30
31
32
33
34
35
36
37
38
39
40
41
42
43
44
45
46
47
48
49
50
51
52
53
54
55
56
57
58
59
60
61
62
63
64
65
66
67
68
69
70
71
72
73
74
75
76
77
78
79
80
81
82
83
84
85
86
87
88
89
90
91
92
93
94
95
96
97
98
99
100
101
102
103
104
105
106
107
108
109
110
111
112
113
114
115
116
117
118
119
120
121
122
123
124
125
126
127
128
129
130
131
132
133
134
135
136
137
138
139
140
141
142
143
144
145
146
147
148
149
150
151
152
153
154
155
156
157
158
159
160
161
162
163
164
165
166
167
168
169
170
171
172
173
174
175
176
177
178
179
180
181
182
183
184
185
186
187
188
189
190
191
192
193
194
195
196
197
198
199
200
201
202
203
204
205
206
207
208
209
210
211
212
213
214
215
216
217
218
219
220
221
222
223
224
225
226
227
228
229
230
231
232
233
234
235
236
237
238
239
240
241
242
243
244
245
246
247
248
249
250
251
252
253
254
255
256
257
258
package ipc
import (
"bufio"
"bytes"
"context"
"encoding/json"
"errors"
"fmt"
"net"
"os"
"path/filepath"
"strings"
"syscall"
"time"
"github.com/thinkwright/threadmark/internal/core"
)
const (
SocketDirMode os.FileMode = 0o700
SocketFileMode os.FileMode = 0o600
DefaultSocket = "daemon.sock"
MaxEventBytes = 1 << 20
)
var sendRetryDelays = []time.Duration{
25 * time.Millisecond,
50 * time.Millisecond,
100 * time.Millisecond,
200 * time.Millisecond,
}
type Handler interface {
HandleEvent(context.Context, core.Event) error
}
type HandlerFunc func(context.Context, core.Event) error
func (fn HandlerFunc) HandleEvent(ctx context.Context, event core.Event) error {
return fn(ctx, event)
}
type Server struct {
SocketPath string
Handler Handler
ErrorHandler func(error)
}
func DefaultSocketPath() (string, error) {
home, err := os.UserHomeDir()
if err != nil {
return "", fmt.Errorf("find user home: %w", err)
}
return filepath.Join(home, ".threadmark", DefaultSocket), nil
}
func (s Server) ListenAndServe(ctx context.Context) error {
socketPath, err := s.socketPath()
if err != nil {
return err
}
if err := prepareSocket(socketPath); err != nil {
return err
}
listener, err := net.Listen("unix", socketPath)
if err != nil {
return fmt.Errorf("listen on unix socket: %w", err)
}
defer listener.Close()
defer os.Remove(socketPath)
if err := os.Chmod(socketPath, SocketFileMode); err != nil {
return fmt.Errorf("set socket permissions: %w", err)
}
go func() {
<-ctx.Done()
_ = listener.Close()
}()
for {
conn, err := listener.Accept()
if err != nil {
if ctx.Err() != nil {
return nil
}
return fmt.Errorf("accept unix socket connection: %w", err)
}
go s.handleConn(ctx, conn)
}
}
func (s Server) handleConn(ctx context.Context, conn net.Conn) {
defer conn.Close()
scanner := bufio.NewScanner(conn)
scanner.Buffer(make([]byte, 0, 64*1024), MaxEventBytes)
for scanner.Scan() {
line := scanner.Bytes()
if len(bytes.TrimSpace(line)) == 0 {
continue
}
if err := s.handleLine(ctx, line); err != nil {
s.handleError(err)
}
}
if err := scanner.Err(); err != nil {
if strings.Contains(err.Error(), "token too long") {
s.handleError(fmt.Errorf("read event: event line exceeds %d bytes", MaxEventBytes))
return
}
s.handleError(fmt.Errorf("read event: %w", err))
}
}
func (s Server) handleLine(ctx context.Context, line []byte) error {
var event core.Event
if err := json.Unmarshal(bytes.TrimSpace(line), &event); err != nil {
return fmt.Errorf("decode event: %w", err)
}
if err := event.Validate(); err != nil {
return fmt.Errorf("validate event: %w", err)
}
if s.Handler == nil {
return nil
}
return s.Handler.HandleEvent(ctx, event)
}
func (s Server) handleError(err error) {
if err != nil && s.ErrorHandler != nil {
s.ErrorHandler(err)
}
}
func Send(ctx context.Context, socketPath string, event core.Event) error {
return SendMany(ctx, socketPath, []core.Event{event})
}
func SendMany(ctx context.Context, socketPath string, events []core.Event) error {
if strings.TrimSpace(socketPath) == "" {
var err error
socketPath, err = DefaultSocketPath()
if err != nil {
return err
}
}
if len(events) == 0 {
return nil
}
var payload []byte
for idx, event := range events {
event = event.Normalize()
if err := event.Validate(); err != nil {
return fmt.Errorf("validate event %d: %w", idx, err)
}
encoded, err := json.Marshal(event)
if err != nil {
return fmt.Errorf("encode event %d: %w", idx, err)
}
payload = append(payload, encoded...)
payload = append(payload, '\n')
}
var lastErr error
for attempt := 0; ; attempt++ {
if err := sendPayload(ctx, socketPath, payload); err != nil {
lastErr = err
} else {
return nil
}
if ctx.Err() != nil {
return lastErr
}
if !isSocketRetryable(lastErr) {
return lastErr
}
if attempt >= len(sendRetryDelays) {
return lastErr
}
timer := time.NewTimer(sendRetryDelays[attempt])
select {
case <-ctx.Done():
timer.Stop()
return lastErr
case <-timer.C:
}
}
}
func sendPayload(ctx context.Context, socketPath string, payload []byte) error {
var dialer net.Dialer
conn, err := dialer.DialContext(ctx, "unix", socketPath)
if err != nil {
return fmt.Errorf("dial threadmark daemon: %w", err)
}
defer conn.Close()
if deadline, ok := ctx.Deadline(); ok {
_ = conn.SetDeadline(deadline)
}
if _, err := conn.Write(payload); err != nil {
return fmt.Errorf("write event: %w", err)
}
return nil
}
func isSocketRetryable(err error) bool {
return errors.Is(err, os.ErrNotExist) ||
errors.Is(err, syscall.ECONNREFUSED) ||
errors.Is(err, syscall.ECONNRESET) ||
errors.Is(err, syscall.EPIPE)
}
func prepareSocket(socketPath string) error {
if strings.TrimSpace(socketPath) == "" {
return errors.New("socket path is required")
}
if err := os.MkdirAll(filepath.Dir(socketPath), SocketDirMode); err != nil {
return fmt.Errorf("create socket directory: %w", err)
}
if err := os.Chmod(filepath.Dir(socketPath), SocketDirMode); err != nil {
return fmt.Errorf("set socket directory permissions: %w", err)
}
info, err := os.Lstat(socketPath)
if errors.Is(err, os.ErrNotExist) {
return nil
}
if err != nil {
return fmt.Errorf("stat socket path: %w", err)
}
if info.Mode()&os.ModeSocket == 0 {
return fmt.Errorf("socket path exists and is not a unix socket: %s", socketPath)
}
conn, err := net.DialTimeout("unix", socketPath, 100*time.Millisecond)
if err == nil {
_ = conn.Close()
return fmt.Errorf("threadmark daemon already listening on %s", socketPath)
}
if err := os.Remove(socketPath); err != nil {
return fmt.Errorf("remove stale socket: %w", err)
}
return nil
}
func (s Server) socketPath() (string, error) {
if strings.TrimSpace(s.SocketPath) != "" {
return s.SocketPath, nil
}
return DefaultSocketPath()
}