-
Notifications
You must be signed in to change notification settings - Fork 1.5k
Expand file tree
/
Copy pathserdes.go
More file actions
140 lines (129 loc) · 3.49 KB
/
Copy pathserdes.go
File metadata and controls
140 lines (129 loc) · 3.49 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
package vpn
import (
"context"
"encoding/binary"
"io"
"sync"
"google.golang.org/protobuf/proto"
"cdr.dev/slog/v3"
)
// MaxLength is the largest possible CoderVPN Protocol message size. This is set
// so that a misbehaving peer can't cause us to allocate a huge amount of memory.
const MaxLength = 0x1000000 // 16MiB
// serdes SERializes and DESerializes protobuf messages to and from the conn.
type serdes[S rpcMessage, R receivableRPCMessage[RR], RR any] struct {
ctx context.Context
logger slog.Logger
conn io.ReadWriteCloser
sendCh <-chan S
recvCh chan<- R
closeOnce sync.Once
wg sync.WaitGroup
}
func (s *serdes[_, R, RR]) recvLoop() {
s.logger.Debug(s.ctx, "starting recvLoop")
defer s.closeIdempotent()
defer close(s.recvCh)
for {
var length uint32
if err := binary.Read(s.conn, binary.BigEndian, &length); err != nil {
s.logger.Debug(s.ctx, "failed to read length", slog.Error(err))
return
}
if length > MaxLength {
s.logger.Critical(s.ctx, "message length exceeds max",
slog.F("length", length))
return
}
s.logger.Debug(s.ctx, "about to read message", slog.F("length", length))
mb := make([]byte, length)
if n, err := io.ReadFull(s.conn, mb); err != nil {
s.logger.Debug(s.ctx, "failed to read message",
slog.Error(err),
slog.F("num_bytes_read", n))
return
}
msg := R(new(RR))
if err := proto.Unmarshal(mb, msg); err != nil {
s.logger.Critical(s.ctx, "failed to unmarshal message", slog.Error(err))
return
}
select {
case s.recvCh <- msg:
s.logger.Debug(s.ctx, "passed received message to speaker")
case <-s.ctx.Done():
s.logger.Debug(s.ctx, "recvLoop canceled", slog.Error(s.ctx.Err()))
}
}
}
func (s *serdes[S, _, _]) sendLoop() {
s.logger.Debug(s.ctx, "starting sendLoop")
defer s.closeIdempotent()
for {
select {
case <-s.ctx.Done():
s.logger.Debug(s.ctx, "sendLoop canceled", slog.Error(s.ctx.Err()))
return
case msg, ok := <-s.sendCh:
if !ok {
s.logger.Debug(s.ctx, "sendCh closed")
return
}
mb, err := proto.Marshal(msg)
if err != nil {
s.logger.Critical(s.ctx, "failed to marshal message", slog.Error(err))
return
}
// #nosec G115 - Safe conversion as protobuf message length is expected to be within uint32 range
if err := binary.Write(s.conn, binary.BigEndian, uint32(len(mb))); err != nil {
s.logger.Debug(s.ctx, "failed to write length", slog.Error(err))
return
}
if _, err := s.conn.Write(mb); err != nil {
s.logger.Debug(s.ctx, "failed to write message", slog.Error(err))
return
}
}
}
}
func (s *serdes[_, _, _]) closeIdempotent() {
s.closeOnce.Do(func() {
if err := s.conn.Close(); err != nil {
s.logger.Error(s.ctx, "failed to close connection", slog.Error(err))
} else {
s.logger.Info(s.ctx, "closed connection")
}
})
}
// Close closes the serdes
// nolint: revive
func (s *serdes[_, _, _]) Close() error {
s.closeIdempotent()
s.wg.Wait()
return nil
}
// start starts the goroutines that serialize and deserialize to the conn.
// nolint: revive
func (s *serdes[_, _, _]) start() {
s.wg.Add(2)
go func() {
defer s.wg.Done()
s.recvLoop()
}()
go func() {
defer s.wg.Done()
s.sendLoop()
}()
}
func newSerdes[S rpcMessage, R receivableRPCMessage[RR], RR any](
ctx context.Context, logger slog.Logger, conn io.ReadWriteCloser,
sendCh <-chan S, recvCh chan<- R,
) *serdes[S, R, RR] {
return &serdes[S, R, RR]{
ctx: ctx,
logger: logger.Named("serdes"),
conn: conn,
sendCh: sendCh,
recvCh: recvCh,
}
}