Skip to content
File

Blob: overlay/transport.go

go500 lines
1package overlay
2 
3import (
4 "context"
5 "crypto/tls"
6 "crypto/x509"
7 "errors"
8 "fmt"
9 "net"
10 "strconv"
11 "strings"
12 "sync"
13 "time"
14 
15 "go.miragespace.co/specter/spec/pki"
16 "go.miragespace.co/specter/spec/protocol"
17 "go.miragespace.co/specter/spec/rpc"
18 "go.miragespace.co/specter/spec/transport"
19 "go.miragespace.co/specter/spec/transport/q"
20 "go.miragespace.co/specter/util/atomic"
21 "go.miragespace.co/specter/util/bufconn"
22 
23 "github.com/avast/retry-go/v5"
24 "github.com/quic-go/quic-go"
25 "github.com/zhangyunhao116/skipmap"
26 uberAtomic "go.uber.org/atomic"
27 "go.uber.org/zap"
28)
29 
30var (
31 _ transport.Transport = (*QUIC)(nil)
32 _ transport.ClientTransport = (*QUIC)(nil)
33)
34 
35var builderPool = sync.Pool{
36 New: func() any {
37 return &strings.Builder{}
38 },
39}
40 
41func NewQUIC(conf TransportConfig) *QUIC {
42 if conf.VirtualTransport && conf.UseCertificateIdentity {
43 panic("cannot enable UseCertificateIdentity and VirtualTransport in the same transport")
44 }
45 return &QUIC{
46 TransportConfig: conf,
47 
48 cachedConnections: skipmap.NewString[*nodeConnection](),
49 cachedMutex: atomic.NewKeyedRWMutex(),
50 
51 streamChan: make(chan *transport.StreamDelegate, 32),
52 dgramChan: make(chan *transport.DatagramDelegate, 32),
53 
54 rttChan: make(chan *transport.DatagramDelegate, 8),
55 rttMap: skipmap.NewString[*skipmap.Uint64Map[int64]](),
56 
57 started: uberAtomic.NewBool(false),
58 closed: uberAtomic.NewBool(false),
59 }
60}
61 
62func (t *QUIC) makeCachedKey(peer *protocol.Node) string {
63 sb := builderPool.Get().(*strings.Builder)
64 defer builderPool.Put(sb)
65 defer sb.Reset()
66 
67 sb.WriteString(peer.GetAddress())
68 sb.WriteString("/")
69 
70 if peer.GetUnknown() {
71 sb.WriteString("-1")
72 } else {
73 if t.VirtualTransport {
74 sb.WriteString("PHY")
75 } else {
76 sb.WriteString(strconv.FormatUint(peer.GetId(), 10))
77 }
78 }
79 return sb.String()
80}
81 
82func (t *QUIC) getCachedConnection(ctx context.Context, peer *protocol.Node) (*quic.Conn, error) {
83 qKey := t.makeCachedKey(peer)
84 
85 if t.Endpoint.GetAddress() == peer.GetAddress() {
86 return nil, fmt.Errorf("creating a new QUIC connection to the ourselves is not allowed")
87 }
88 
89 retrier := retry.NewWithData[*quic.Conn](
90 retry.Attempts(2),
91 retry.Context(ctx),
92 retry.LastErrorOnly(true),
93 retry.OnRetry(func(n uint, err error) {
94 t.Logger.Info("Potential connection reuse conflict, retrying to get previously cached connection", zap.Object("peer", peer), zap.Error(err))
95 }),
96 retry.RetryIf(func(err error) bool {
97 return strings.Contains(err.Error(), reuseErrorState)
98 }),
99 )
100 q, err := retrier.Do(func() (*quic.Conn, error) {
101 rUnlock := t.cachedMutex.RLock(qKey)
102 if cached, ok := t.cachedConnections.Load(qKey); ok {
103 rUnlock()
104 return cached.quic, nil
105 }
106 rUnlock()
107 
108 if peer.GetRendezvous() || peer.GetAddress() == "" {
109 return nil, transport.ErrNoDirect
110 }
111 
112 t.Logger.Debug("Creating new QUIC connection", zap.Object("peer", peer))
113 
114 dialCtx, dialCancel := context.WithTimeout(ctx, transport.ConnectTimeout)
115 defer dialCancel()
116 
117 peerAddr := peer.GetAddress()
118 
119 addr, err := net.ResolveUDPAddr("udp", peerAddr)
120 if err != nil {
121 return nil, err
122 }
123 
124 cfg := t.ClientTLS.Clone()
125 if cert, ok := t.clientCert.Load().(tls.Certificate); ok {
126 cfg.Certificates = []tls.Certificate{cert}
127 }
128 
129 if cfg.ServerName == "" {
130 host, _, err := net.SplitHostPort(peerAddr)
131 if err != nil {
132 host = peerAddr
133 }
134 cfg.ServerName = host
135 }
136 
137 newQ, err := t.QuicTransport.DialEarly(dialCtx, addr, cfg, quicConfig)
138 if err != nil {
139 return nil, err
140 }
141 
142 return t.handleOutgoing(ctx, newQ)
143 })
144 if err != nil {
145 if err != transport.ErrNoDirect {
146 t.Logger.Error("Failed to establish connection", zap.Object("peer", peer), zap.Error(err))
147 }
148 return nil, err
149 }
150 
151 t.background(ctx)
152 
153 return q, nil
154}
155 
156func (t *QUIC) WithClientCertificate(cert tls.Certificate) error {
157 if len(cert.Certificate) == 0 {
158 return transport.ErrNoCertificate
159 }
160 if cert.PrivateKey == nil {
161 return transport.ErrNoCertificate
162 }
163 parsed, err := x509.ParseCertificate(cert.Certificate[0])
164 if err != nil {
165 return err
166 }
167 identity, err := pki.ExtractCertificateIdentity(parsed)
168 if err != nil {
169 return err
170 }
171 
172 t.Logger.Debug("Using client certificate", zap.Object("identity", identity))
173 t.Endpoint.Id = identity.ID
174 t.clientCert.Store(cert)
175 
176 return nil
177}
178 
179func (t *QUIC) Identity() *protocol.Node {
180 return t.Endpoint
181}
182 
183func (t *QUIC) DialStream(ctx context.Context, peer *protocol.Node, kind protocol.Stream_Type) (net.Conn, error) {
184 if t.closed.Load() {
185 return nil, transport.ErrClosed
186 }
187 
188 if peer.GetAddress() == t.Endpoint.GetAddress() && t.VirtualTransport {
189 c1, c2 := bufconn.BufferedPipe(8192)
190 t.streamChan <- &transport.StreamDelegate{
191 Identity: &protocol.Node{
192 Address: t.Endpoint.GetAddress(),
193 Id: peer.GetId(),
194 },
195 Conn: c2,
196 Kind: kind,
197 }
198 return c1, nil
199 }
200 
201 q, err := t.getCachedConnection(ctx, peer)
202 if err != nil {
203 return nil, fmt.Errorf("creating quic connection: %w", err)
204 }
205 
206 var rr protocol.Stream
207 if t.VirtualTransport {
208 rr = protocol.Stream{
209 Type: kind,
210 Target: &protocol.Node{
211 Id: peer.GetId(),
212 },
213 }
214 } else {
215 rr = protocol.Stream{
216 Type: kind,
217 }
218 }
219 return openStream(q, &rr)
220}
221 
222func openStream(q *quic.Conn, rr *protocol.Stream) (net.Conn, error) {
223 stream, err := q.OpenStream()
224 if err != nil {
225 return nil, err
226 }
227 conn := WrapQuicConnection(stream, q)
228 stream.SetDeadline(time.Now().Add(quicConfig.HandshakeIdleTimeout))
229 if err := rpc.Send(stream, rr); err != nil {
230 conn.Close()
231 return nil, err
232 }
233 stream.SetDeadline(time.Time{})
234 return conn, nil
235}
236 
237func (t *QUIC) AcceptStream() <-chan *transport.StreamDelegate {
238 return t.streamChan
239}
240 
241func (t *QUIC) SupportDatagram() bool {
242 return quicConfig.EnableDatagrams
243}
244 
245func (t *QUIC) ReceiveDatagram() <-chan *transport.DatagramDelegate {
246 return t.dgramChan
247}
248 
249func (t *QUIC) SendDatagram(peer *protocol.Node, buf []byte) error {
250 qKey := t.makeCachedKey(peer)
251 if r, ok := t.cachedConnections.Load(qKey); ok {
252 data := &protocol.Datagram{
253 Type: protocol.Datagram_DATA,
254 Data: buf,
255 }
256 b, err := data.MarshalVT()
257 if err != nil {
258 return err
259 }
260 return r.quic.SendDatagram(b)
261 }
262 return transport.ErrNoDirect
263}
264 
265func (t *QUIC) handleIncoming(ctx context.Context, q *quic.Conn) (*quic.Conn, error) {
266 openCtx, openCancel := context.WithTimeout(ctx, quicConfig.HandshakeIdleTimeout)
267 defer openCancel()
268 
269 stream, err := q.OpenStreamSync(openCtx)
270 if err != nil {
271 return nil, err
272 }
273 defer WrapQuicConnection(stream, q).Close()
274 
275 c, reused, err := t.reuseConnection(ctx, q, stream, directionIncoming)
276 if err != nil {
277 return nil, err
278 }
279 
280 if !reused {
281 t.handlePeer(ctx, c.quic, c.peer, directionIncoming)
282 }
283 
284 return c.quic, nil
285}
286 
287func (t *QUIC) handleOutgoing(ctx context.Context, q *quic.Conn) (*quic.Conn, error) {
288 openCtx, openCancel := context.WithTimeout(ctx, quicConfig.HandshakeIdleTimeout)
289 defer openCancel()
290 
291 stream, err := q.AcceptStream(openCtx)
292 if err != nil {
293 return nil, err
294 }
295 defer WrapQuicConnection(stream, q).Close()
296 
297 c, reused, err := t.reuseConnection(ctx, q, stream, directionOutgoing)
298 if err != nil {
299 return nil, err
300 }
301 
302 if !reused {
303 t.handlePeer(ctx, c.quic, c.peer, directionOutgoing)
304 }
305 
306 return c.quic, nil
307}
308 
309func (t *QUIC) handlePeer(ctx context.Context, q *quic.Conn, peer *protocol.Node, dir direction) {
310 l := t.Logger.With(
311 zap.String("remote", q.RemoteAddr().String()),
312 zap.Object("peer", peer),
313 zap.String("direction", dir.String()),
314 zap.String("key", t.makeCachedKey(peer)),
315 )
316 l.Debug("Starting goroutines to handle streams and datagrams")
317 go t.handleConnection(ctx, q, peer)
318 go t.handleDatagram(ctx, q, peer)
319 if t.RTTRecorder != nil {
320 go t.sendRTTSyn(ctx, q, peer)
321 }
322 go func(q *quic.Conn) {
323 <-q.Context().Done()
324 l.Debug("Connection with peer closed", zap.Error(q.Context().Err()))
325 t.reapPeer(q, peer)
326 }(q)
327}
328 
329func (t *QUIC) background(ctx context.Context) {
330 if !t.started.CompareAndSwap(false, true) {
331 return
332 }
333 go t.reaper(ctx)
334 go t.handleRTTAck(ctx)
335}
336 
337func (t *QUIC) AcceptWithListener(ctx context.Context, listener q.Listener) error {
338 t.Logger.Info("Accepting connections", zap.String("listen", listener.Addr().String()))
339 t.background(ctx)
340 for {
341 q, err := listener.Accept(ctx)
342 if err != nil {
343 return err
344 }
345 go func(q *quic.Conn) {
346 if _, err := t.handleIncoming(ctx, q); err != nil {
347 if !strings.Contains(err.Error(), reuseErrorState) {
348 t.Logger.Error("Incoming connection reuse error", zap.String("endpoint", q.RemoteAddr().String()), zap.Error(err))
349 }
350 // TODO: figure out a better way to ensure that the peer received cache status before closing
351 time.Sleep(time.Second)
352 q.CloseWithError(406, err.Error())
353 }
354 }(q)
355 }
356}
357 
358func (t *QUIC) handleDatagram(ctx context.Context, q *quic.Conn, peer *protocol.Node) {
359 logger := t.Logger.With(zap.String("endpoint", q.RemoteAddr().String()), zap.Object("peer", peer))
360 for {
361 b, err := q.ReceiveDatagram(ctx)
362 if err != nil {
363 if !errors.Is(err, net.ErrClosed) {
364 logger.Error("error receiving datagram", zap.Error(err))
365 }
366 return
367 }
368 data := &protocol.Datagram{}
369 if err := data.UnmarshalVT(b); err != nil {
370 logger.Error("error decoding datagram to proto", zap.Error(err))
371 continue
372 }
373 switch data.GetType() {
374 case protocol.Datagram_ALIVE:
375 case protocol.Datagram_RTT_SYN:
376 data.Type = protocol.Datagram_RTT_ACK
377 buf, err := data.MarshalVT()
378 if err != nil {
379 logger.Error("error encoding rtt ack datagram to proto", zap.Error(err))
380 continue
381 }
382 if err := q.SendDatagram(buf); err != nil {
383 logger.Error("error sending rtt ack datagram", zap.Error(err))
384 continue
385 }
386 case protocol.Datagram_RTT_ACK:
387 select {
388 case t.rttChan <- &transport.DatagramDelegate{Buffer: data.GetData(), Identity: peer}:
389 default:
390 logger.Warn("rtt ack buffer full, dropping datagram")
391 }
392 case protocol.Datagram_DATA:
393 select {
394 case t.dgramChan <- &transport.DatagramDelegate{Buffer: data.GetData(), Identity: peer}:
395 default:
396 logger.Warn("data buffer full, dropping datagram")
397 }
398 default:
399 logger.Warn("unknown datagram type: %s", zap.String("type", data.GetType().String()))
400 }
401 }
402}
403 
404func (t *QUIC) handleConnection(ctx context.Context, q *quic.Conn, peer *protocol.Node) {
405 for {
406 stream, err := q.AcceptStream(ctx)
407 if err != nil {
408 if !errors.Is(err, net.ErrClosed) {
409 t.Logger.Error("Error accepting new stream from peer", zap.Object("peer", peer), zap.String("remote", q.RemoteAddr().String()), zap.Error(err))
410 }
411 return
412 }
413 go t.streamHandler(q, stream, peer)
414 }
415}
416 
417func (t *QUIC) streamHandler(q *quic.Conn, stream *quic.Stream, peer *protocol.Node) {
418 l := t.Logger.With(zap.Object("peer", peer))
419 conn := WrapQuicConnection(stream, q)
420 
421 var err error
422 defer func() {
423 if err != nil {
424 l.Error("error handshaking on new stream", zap.Error(err))
425 conn.Close()
426 }
427 }()
428 
429 rr := protocol.Stream{}
430 conn.SetDeadline(time.Now().Add(quicConfig.HandshakeIdleTimeout))
431 err = rpc.BoundedReceive(conn, &rr, 16)
432 if err != nil {
433 l.Error("Failed to receive stream handshake", zap.Error(err))
434 conn.Close()
435 return
436 }
437 conn.SetDeadline(time.Time{})
438 
439 if rr.GetType() == protocol.Stream_UNKNOWN_TYPE {
440 l.Warn("Received stream with unknown type")
441 conn.Close()
442 return
443 }
444 
445 var identity *protocol.Node
446 if t.VirtualTransport {
447 identity = &protocol.Node{
448 Id: rr.GetTarget().GetId(),
449 Address: peer.GetAddress(),
450 }
451 } else {
452 identity = peer
453 }
454 
455 delegation := &transport.StreamDelegate{
456 Identity: identity,
457 Conn: conn,
458 Kind: rr.GetType(),
459 }
460 
461 if t.UseCertificateIdentity {
462 chain := q.ConnectionState().TLS.VerifiedChains
463 delegation.Certificate = chain[0][0]
464 }
465 
466 select {
467 case t.streamChan <- delegation:
468 default:
469 l.Warn("Stream channel full, dropping incoming stream",
470 zap.String("kind", rr.GetType().String()),
471 )
472 conn.Close()
473 }
474}
475 
476func (t *QUIC) ListConnected() []transport.ConnectedPeer {
477 nodes := make([]transport.ConnectedPeer, 0)
478 t.cachedConnections.Range(func(key string, value *nodeConnection) bool {
479 nodes = append(nodes, transport.ConnectedPeer{
480 Identity: value.peer,
481 Addr: value.quic.RemoteAddr(),
482 Version: value.version,
483 Physical: attachment{value.quic},
484 })
485 return true
486 })
487 return nodes
488}
489 
490func (t *QUIC) Stop() {
491 if !t.closed.CompareAndSwap(false, true) {
492 return
493 }
494 t.started.Store(false)
495 t.cachedConnections.Range(func(key string, value *nodeConnection) bool {
496 value.quic.CloseWithError(0, "Transport closed")
497 return true
498 })
499}