master
go 253 lines 5.03 KB
Raw
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.Go(func() {
68 t.handleConnection(conn)
69 })
70 }
71 }
72 }
73
74 func (t *tcpServer) handleConnection(conn net.Conn) {
75 defer func() { _ = conn.Close() }()
76
77 if err := conn.SetDeadline(time.Now().Add(time.Second)); err != nil {
78 return
79 }
80
81 rw := bufio.NewReadWriter(bufio.NewReader(conn), bufio.NewWriter(conn))
82
83 if _, err := rw.ReadString('\n'); err != nil {
84 writeResponse(rw, fmt.Sprintf("failed to read input: %v\n", err))
85 } else {
86 writeResponse(rw, "pong\n")
87 }
88 }
89
90 func newUDPServer(addr string) *udpServer {
91 ctx, cancel := context.WithCancel(context.Background())
92 _, addr = parseAddress(addr)
93 return &udpServer{
94 addr: addr,
95 ctx: ctx,
96 cancel: cancel,
97 }
98 }
99
100 type udpServer struct {
101 addr string
102 conn *net.UDPConn
103 ctx context.Context
104 cancel context.CancelFunc
105 }
106
107 func (u *udpServer) Run() error {
108 addr, err := net.ResolveUDPAddr("udp", u.addr)
109 if err != nil {
110 return fmt.Errorf("failed to resolve UDP address: %w", err)
111 }
112
113 u.conn, err = net.ListenUDP("udp", addr)
114 if err != nil {
115 return fmt.Errorf("failed to start UDP server: %w", err)
116 }
117
118 return u.handleConnections()
119 }
120
121 func (u *udpServer) Close() (err error) {
122 u.cancel()
123 if u.conn != nil {
124 if err := u.conn.Close(); err != nil {
125 return fmt.Errorf("failed to close UDP server: %w", err)
126 }
127 }
128 return nil
129 }
130
131 func (u *udpServer) handleConnections() error {
132 buffer := make([]byte, 8192)
133 for {
134 select {
135 case <-u.ctx.Done():
136 return nil
137 default:
138 if err := u.conn.SetReadDeadline(time.Now().Add(time.Second)); err != nil {
139 continue
140 }
141
142 _, addr, err := u.conn.ReadFromUDP(buffer[0:])
143 if err != nil {
144 if !errors.Is(err, os.ErrDeadlineExceeded) {
145 return fmt.Errorf("failed to read UDP packet: %w", err)
146 }
147 continue
148 }
149
150 if _, err := u.conn.WriteToUDP([]byte("pong\n"), addr); err != nil {
151 return fmt.Errorf("failed to write UDP response: %w", err)
152 }
153 }
154 }
155 }
156
157 func newUnixServer(addr string) *unixServer {
158 ctx, cancel := context.WithCancel(context.Background())
159 _, addr = parseAddress(addr)
160 return &unixServer{
161 addr: addr,
162 ctx: ctx,
163 cancel: cancel,
164 }
165 }
166
167 type unixServer struct {
168 addr string
169 listener *net.UnixListener
170 wg sync.WaitGroup
171 ctx context.Context
172 cancel context.CancelFunc
173 }
174
175 func (u *unixServer) Run() error {
176 if err := os.Remove(u.addr); err != nil && !os.IsNotExist(err) {
177 return fmt.Errorf("failed to clean up existing socket: %w", err)
178 }
179
180 addr, err := net.ResolveUnixAddr("unix", u.addr)
181 if err != nil {
182 return fmt.Errorf("failed to resolve Unix address: %w", err)
183 }
184
185 u.listener, err = net.ListenUnix("unix", addr)
186 if err != nil {
187 return fmt.Errorf("failed to start Unix server: %w", err)
188 }
189
190 return u.handleConnections()
191 }
192
193 func (u *unixServer) Close() error {
194 u.cancel()
195
196 if u.listener != nil {
197 if err := u.listener.Close(); err != nil {
198 return fmt.Errorf("failed to close Unix server: %w", err)
199 }
200 }
201
202 u.wg.Wait()
203 _ = os.Remove(u.addr)
204
205 return nil
206 }
207
208 func (u *unixServer) handleConnections() error {
209 for {
210 select {
211 case <-u.ctx.Done():
212 return nil
213 default:
214 if err := u.listener.SetDeadline(time.Now().Add(time.Second)); err != nil {
215 continue
216 }
217
218 conn, err := u.listener.AcceptUnix()
219 if err != nil {
220 if !errors.Is(err, os.ErrDeadlineExceeded) {
221 return err
222 }
223 continue
224 }
225
226 u.wg.Go(func() {
227 u.handleConnection(conn)
228 })
229 }
230 }
231 }
232
233 func (u *unixServer) handleConnection(conn net.Conn) {
234 defer func() { _ = conn.Close() }()
235
236 if err := conn.SetDeadline(time.Now().Add(time.Second)); err != nil {
237 return
238 }
239
240 rw := bufio.NewReadWriter(bufio.NewReader(conn), bufio.NewWriter(conn))
241
242 if _, err := rw.ReadString('\n'); err != nil {
243 writeResponse(rw, fmt.Sprintf("failed to read input: %v\n", err))
244 } else {
245 writeResponse(rw, "pong\n")
246 }
247
248 }
249
250 func writeResponse(rw *bufio.ReadWriter, response string) {
251 _, _ = rw.WriteString(response)
252 _ = rw.Flush()
253 }