File
Blob: overlay/transport.go
| 1 | package overlay |
| 2 | |
| 3 | import ( |
| 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 | |
| 30 | var ( |
| 31 | _ transport.Transport = (*QUIC)(nil) |
| 32 | _ transport.ClientTransport = (*QUIC)(nil) |
| 33 | ) |
| 34 | |
| 35 | var builderPool = sync.Pool{ |
| 36 | New: func() any { |
| 37 | return &strings.Builder{} |
| 38 | }, |
| 39 | } |
| 40 | |
| 41 | func 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 | |
| 62 | func (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 | |
| 82 | func (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 | |
| 156 | func (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 | |
| 179 | func (t *QUIC) Identity() *protocol.Node { |
| 180 | return t.Endpoint |
| 181 | } |
| 182 | |
| 183 | func (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 | |
| 222 | func 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 | |
| 237 | func (t *QUIC) AcceptStream() <-chan *transport.StreamDelegate { |
| 238 | return t.streamChan |
| 239 | } |
| 240 | |
| 241 | func (t *QUIC) SupportDatagram() bool { |
| 242 | return quicConfig.EnableDatagrams |
| 243 | } |
| 244 | |
| 245 | func (t *QUIC) ReceiveDatagram() <-chan *transport.DatagramDelegate { |
| 246 | return t.dgramChan |
| 247 | } |
| 248 | |
| 249 | func (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 | |
| 265 | func (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 | |
| 287 | func (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 | |
| 309 | func (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 | |
| 329 | func (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 | |
| 337 | func (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 | |
| 358 | func (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 | |
| 404 | func (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 | |
| 417 | func (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 | |
| 476 | func (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 | |
| 490 | func (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 | } |