main
go 961 lines 24.7 KB
Raw
1 package sdk
2
3 import (
4 "bufio"
5 "bytes"
6 "context"
7 "crypto/tls"
8 "encoding/base64"
9 "errors"
10 "fmt"
11 "io"
12 "net"
13 "net/http"
14 "net/url"
15 "strings"
16 "sync"
17 "time"
18
19 "github.com/quic-go/quic-go"
20 "github.com/rs/zerolog/log"
21
22 "github.com/gosuda/portal-tunnel/v2/portal/discovery"
23 "github.com/gosuda/portal-tunnel/v2/portal/identity"
24 "github.com/gosuda/portal-tunnel/v2/portal/keyless"
25 "github.com/gosuda/portal-tunnel/v2/portal/transport"
26 "github.com/gosuda/portal-tunnel/v2/types"
27 "github.com/gosuda/portal-tunnel/v2/utils"
28 )
29
30 type listenerConfig struct {
31 Identity types.Identity
32 UDPEnabled bool
33 TCPEnabled bool
34 BanMITM bool
35 Metadata func() types.LeaseMetadata
36 DialTimeout time.Duration
37 RequestTimeout time.Duration
38 HandshakeTimeout time.Duration
39 LeaseTTL time.Duration
40 RenewBefore time.Duration
41 ReadyTarget int
42 RetryCount int
43 RetryWait time.Duration
44 relaySet *discovery.RelaySet
45 }
46
47 var errLeaseRefreshRequired = errors.New("lease refresh required")
48
49 type listener struct {
50 cancel context.CancelFunc
51 doneCh <-chan struct{}
52 closeOnce sync.Once
53
54 relayURL *url.URL
55 route discovery.Route
56 metadata func() types.LeaseMetadata
57 identity types.Identity
58 relaySet *discovery.RelaySet
59 udpEnabled bool
60 tcpEnabled bool
61 dialTimeout time.Duration
62 requestTimeout time.Duration
63 readyTarget int
64 retryCount int
65 retryWait time.Duration
66 leaseTTL time.Duration
67 renewBefore time.Duration
68
69 stream *transport.ClientStream
70 datagram *transport.ClientDatagram
71 mitmManager *mitmManager
72
73 httpClient *http.Client
74 httpTransport *http.Transport
75 tlsConfig *tls.Config
76
77 releaseVersion string
78
79 lease *utils.Snapshot[listenerSnapshot]
80 }
81
82 // newListener creates one relay listener and its dedicated relay transport for one relay URL.
83 // Only local config validation fails immediately; relay startup runs in the background until ready.
84 func newListener(ctx context.Context, route discovery.Route, cfg listenerConfig) (*listener, error) {
85 listenerCtx, cancel := context.WithCancel(ctx)
86 readyTarget := utils.IntOrDefault(cfg.ReadyTarget, defaultReadyTarget)
87 leaseTTL := utils.DurationOrDefault(cfg.LeaseTTL, defaultLeaseTTL)
88 dialTimeout := utils.DurationOrDefault(cfg.DialTimeout, defaultDialTimeout)
89 requestTimeout := utils.DurationOrDefault(cfg.RequestTimeout, defaultRequestTimeout)
90 handshakeTimeout := utils.DurationOrDefault(cfg.HandshakeTimeout, defaultHandshakeTimeout)
91 renewBefore := utils.DurationOrDefault(cfg.RenewBefore, defaultRenewBefore)
92 retryWait := utils.DurationOrDefault(cfg.RetryWait, defaultRetryWait)
93
94 normalizedRelayURL, err := utils.NormalizeRelayURL(route.ListenerRelayURL())
95 if err != nil {
96 cancel()
97 return nil, err
98 }
99 relayurl, err := url.Parse(normalizedRelayURL)
100 if err != nil {
101 cancel()
102 return nil, fmt.Errorf("parse relay url: %w", err)
103 }
104 l := &listener{
105 cancel: cancel,
106 doneCh: listenerCtx.Done(),
107 relayURL: relayurl,
108 route: route.WithListenerRelayURL(normalizedRelayURL),
109 metadata: cfg.Metadata,
110 identity: cfg.Identity.Copy(),
111 relaySet: cfg.relaySet,
112 udpEnabled: cfg.UDPEnabled,
113 tcpEnabled: cfg.TCPEnabled,
114 dialTimeout: dialTimeout,
115 requestTimeout: requestTimeout,
116 readyTarget: readyTarget,
117 retryCount: cfg.RetryCount,
118 retryWait: retryWait,
119 leaseTTL: leaseTTL,
120 renewBefore: renewBefore,
121 lease: utils.NewSnapshot(listenerSnapshot{}, listenerSnapshot.snapshot),
122 }
123 l.mitmManager = newMITMManager(listenerCtx, l, cfg.BanMITM)
124 l.stream = transport.NewClientStream(readyTarget, handshakeTimeout)
125 if l.udpEnabled {
126 l.datagram = transport.NewClientDatagram(func(err error) {
127 log.Info().
128 Err(err).
129 Str("component", "sdk-quic-backhaul").
130 Str("address", l.identity.Address).
131 Msg("quic backhaul disconnected; waiting to reconnect")
132 })
133 }
134
135 go l.run(listenerCtx)
136 return l, nil
137 }
138
139 func (l *listener) metadataSnapshot() types.LeaseMetadata {
140 if l.metadata == nil {
141 return types.LeaseMetadata{}
142 }
143 return l.metadata()
144 }
145
146 func (l *listener) run(ctx context.Context) {
147 var retries int
148
149 for {
150 err := l.registerAndConfigure(ctx)
151 switch {
152 case err == nil:
153 case errors.Is(err, context.Canceled), errors.Is(err, net.ErrClosed):
154 return
155 default:
156 if errors.Is(err, errRelayIncompatible) ||
157 errors.Is(err, &types.APIRequestError{Code: types.APIErrorCodeFeatureUnavailable}) ||
158 errors.Is(err, &types.APIRequestError{Code: types.APIErrorCodeTransportMismatch}) ||
159 errors.Is(err, &types.APIRequestError{Code: types.APIErrorCodeHostnameConflict}) ||
160 errors.Is(err, &types.APIRequestError{Code: types.APIErrorCodeIPBanned}) {
161 relayURL := l.relayURL.String()
162 if l.relaySet != nil && relayURL != "" {
163 l.relaySet.UnconfirmRelayURL(relayURL)
164 l.relaySet.RecordActiveFailure(relayURL, 1)
165 }
166 log.Error().
167 Err(err).
168 Str("relay_url", relayURL).
169 Str("address", l.identity.Address).
170 Msg("lease registration failed; closing listener")
171 _ = l.Close()
172 return
173 }
174 retries++
175 if !l.waitRetry(ctx, "lease registration", err, retries, 0) {
176 _ = l.Close()
177 return
178 }
179 continue
180 }
181
182 retries = 0
183 publicURL := ""
184 if lease, ok := l.leaseSnapshot(); ok {
185 publicURL = l.publicURLForLease(lease)
186 }
187 event := log.Info().Str("address", l.identity.Address)
188 if publicURL != "" {
189 event.Msg("service ready at " + publicURL)
190 } else {
191 event.Msg("relay listener registered")
192 }
193
194 err = l.runLease(ctx)
195 if err == nil || errors.Is(err, context.Canceled) || errors.Is(err, net.ErrClosed) {
196 return
197 }
198
199 if errors.Is(err, errLeaseRefreshRequired) {
200 lease := l.clearLease("lease refresh required")
201 if lease != nil && lease.tlsCloser != nil {
202 _ = lease.tlsCloser.Close()
203 }
204 l.resetTransport()
205 relayURL := l.relayURL.String()
206 log.Debug().
207 Err(err).
208 Str("relay_url", relayURL).
209 Str("address", l.identity.Address).
210 Msg("lease refresh required; re-registering")
211 continue
212 }
213
214 relayURL := l.relayURL.String()
215 log.Error().
216 Err(err).
217 Str("relay_url", relayURL).
218 Str("address", l.identity.Address).
219 Msg("listener connection retry budget exhausted; closing listener")
220 _ = l.Close()
221 return
222 }
223 }
224
225 func (l *listener) Close() error {
226 var closeErr error
227 l.closeOnce.Do(func() {
228 if l.cancel != nil {
229 l.cancel()
230 }
231
232 lease := l.clearLease("")
233
234 if l.stream != nil {
235 l.stream.Drain()
236 }
237 if l.datagram != nil {
238 l.datagram.Close()
239 }
240
241 if lease != nil && lease.hostname != "" && l.identity.Key() != "" && lease.accessToken != "" {
242 ctx, cancel := context.WithTimeout(context.Background(), 5*time.Second)
243 closeErr = errors.Join(closeErr, l.unregisterLease(ctx, lease.accessToken, lease.hopRoutes))
244 cancel()
245 }
246 if lease != nil && lease.tlsCloser != nil {
247 closeErr = errors.Join(closeErr, lease.tlsCloser.Close())
248 }
249 l.resetTransport()
250 })
251 return closeErr
252 }
253
254 type listenerSnapshot struct {
255 hostname string
256 echConfigList []byte
257 udpAddr string
258 tcpAddr string
259 accessToken string
260 multihopAccessToken string
261 expiresAt time.Time
262 sniPort int
263 publicURLBase *url.URL
264 tlsConfig *tls.Config
265 tlsCloser io.Closer
266 hopRoutes []types.HopRoute
267 }
268
269 func (s listenerSnapshot) snapshot() listenerSnapshot {
270 s.echConfigList = bytes.Clone(s.echConfigList)
271 s.hopRoutes = append([]types.HopRoute(nil), s.hopRoutes...)
272 if s.publicURLBase != nil {
273 publicURLBase := *s.publicURLBase
274 s.publicURLBase = &publicURLBase
275 }
276 if s.tlsConfig != nil {
277 s.tlsConfig = s.tlsConfig.Clone()
278 }
279 return s
280 }
281
282 func (l *listener) clearLease(reason string) *listenerSnapshot {
283 if l == nil || l.lease == nil {
284 return nil
285 }
286 lease := l.lease.Swap(listenerSnapshot{})
287
288 if l.mitmManager != nil {
289 l.mitmManager.reset()
290 }
291 if l.datagram != nil && reason != "" {
292 l.datagram.Clear(reason)
293 }
294 if lease.accessToken == "" && lease.tlsCloser == nil {
295 return nil
296 }
297 return &lease
298 }
299
300 func (l *listener) leaseSnapshot() (listenerSnapshot, bool) {
301 if l == nil || l.lease == nil {
302 return listenerSnapshot{}, false
303 }
304 lease := l.lease.Load()
305 if lease.accessToken == "" {
306 return listenerSnapshot{}, false
307 }
308 return lease, true
309 }
310
311 func (l *listener) Accept() (net.Conn, error) {
312 if l.stream == nil {
313 return nil, net.ErrClosed
314 }
315 for {
316 conn, err := l.stream.Accept(l.doneCh)
317 if err != nil {
318 return nil, err
319 }
320
321 nextConn, handled, handleErr := l.mitmManager.maybeHandleConn(conn)
322 if handleErr != nil {
323 log.Debug().
324 Err(handleErr).
325 Str("relay_url", l.relayURL.String()).
326 Str("address", l.identity.Address).
327 Msg("mitm self-probe handling failed")
328 }
329 if handled {
330 continue
331 }
332 return &mitmProbeConn{Conn: nextConn, manager: l.mitmManager}, nil
333 }
334 }
335
336 func (l *listener) acceptDatagram() (types.DatagramFrame, error) {
337 if l.datagram == nil {
338 return types.DatagramFrame{}, net.ErrClosed
339 }
340
341 frame, err := l.datagram.Accept(l.doneCh)
342 if err != nil {
343 return types.DatagramFrame{}, err
344 }
345
346 frame.Payload = bytes.Clone(frame.Payload)
347 if lease, ok := l.leaseSnapshot(); ok {
348 frame.UDPAddr = lease.udpAddr
349 }
350 frame.Address = l.identity.Address
351 if l.relayURL != nil {
352 frame.RelayURL = l.relayURL.String()
353 }
354 return frame, nil
355 }
356
357 func (l *listener) sendDatagram(frame types.DatagramFrame) error {
358 if l.datagram == nil {
359 return net.ErrClosed
360 }
361
362 if l.identity.Address == "" {
363 return net.ErrClosed
364 }
365 if frameAddress := strings.TrimSpace(frame.Address); frameAddress != "" && frameAddress != l.identity.Address {
366 return errors.New("datagram frame targets stale address")
367 }
368 return l.datagram.Send(frame.FlowID, frame.Payload)
369 }
370
371 func (l *listener) datagramReady() (string, bool, bool) {
372 if l.datagram == nil {
373 return "", false, false
374 }
375
376 hostname := ""
377 udpAddr := ""
378 if lease, ok := l.leaseSnapshot(); ok {
379 hostname = lease.hostname
380 udpAddr = lease.udpAddr
381 }
382 ready := l.datagram.Connected() && udpAddr != ""
383 closed := false
384 select {
385 case <-l.doneCh:
386 closed = true
387 default:
388 }
389 pending := !ready && !closed && (hostname == "" || udpAddr != "")
390 return udpAddr, ready, pending
391 }
392
393 func (l *listener) publicURLForLease(lease listenerSnapshot) string {
394 baseURL := lease.publicURLBase
395 if baseURL == nil {
396 baseURL = l.relayURL
397 }
398 if baseURL == nil {
399 return ""
400 }
401 if lease.hostname == "" {
402 return ""
403 }
404
405 if baseURL.Scheme == "" {
406 return "https://" + lease.hostname
407 }
408
409 host := lease.hostname
410 sniPort := lease.sniPort
411 scheme := strings.ToLower(strings.TrimSpace(baseURL.Scheme))
412 if (scheme == "https" && sniPort == 443) || (scheme == "http" && sniPort == 80) {
413 sniPort = 0
414 }
415 if sniPort > 0 {
416 host = net.JoinHostPort(lease.hostname, fmt.Sprintf("%d", sniPort))
417 }
418
419 return (&url.URL{
420 Scheme: baseURL.Scheme,
421 Host: host,
422 }).String()
423 }
424
425 func (l *listener) runLease(ctx context.Context) error {
426 lease, ok := l.leaseSnapshot()
427 if !ok || lease.hostname == "" {
428 if ctx.Err() != nil {
429 return ctx.Err()
430 }
431 return errLeaseRefreshRequired
432 }
433 leaseCtx, cancel := context.WithCancel(ctx)
434 defer cancel()
435
436 errCh := make(chan error, max(l.readyTarget, 1)+1)
437 if l.stream != nil && l.readyTarget > 0 {
438 for sessionSlot := range l.readyTarget {
439 sessionSlot++
440 go func() {
441 if err := l.runReverseSessionLoop(leaseCtx, lease.tlsConfig, sessionSlot); err != nil {
442 select {
443 case errCh <- err:
444 case <-leaseCtx.Done():
445 }
446 }
447 }()
448 }
449 }
450 if l.udpEnabled {
451 go l.runDatagramLoop(leaseCtx)
452 }
453 go func() {
454 if err := l.runRenewLoop(leaseCtx); err != nil {
455 select {
456 case errCh <- err:
457 case <-leaseCtx.Done():
458 }
459 }
460 }()
461
462 select {
463 case <-ctx.Done():
464 return ctx.Err()
465 case err := <-errCh:
466 cancel()
467 return err
468 }
469 }
470
471 func (l *listener) runReverseSessionLoop(ctx context.Context, tlsConfig *tls.Config, sessionSlot int) error {
472 if l.stream == nil {
473 return nil
474 }
475
476 var retries int
477 for {
478 conn, err := l.openReverseSession(ctx)
479 if err != nil {
480 if errors.Is(err, context.Canceled) || errors.Is(err, net.ErrClosed) {
481 return nil
482 }
483 retries++
484 if !l.waitRetry(ctx, "reverse session connect", err, retries, sessionSlot) {
485 return err
486 }
487 continue
488 }
489
490 claimed, err := l.stream.RunSession(ctx, conn, tlsConfig)
491 switch {
492 case err == nil:
493 retries = 0
494 case errors.Is(err, context.Canceled), errors.Is(err, net.ErrClosed):
495 return nil
496 case claimed:
497 log.Debug().
498 Err(err).
499 Str("relay_url", l.relayURL.String()).
500 Str("address", l.identity.Address).
501 Int("reverse_session_slot", sessionSlot).
502 Msg("tenant tls handshake failed")
503 retries = 0
504 default:
505 retries++
506 if !l.waitRetry(ctx, "reverse session connect", err, retries, sessionSlot) {
507 return err
508 }
509 }
510 }
511 }
512
513 func (l *listener) runDatagramLoop(ctx context.Context) {
514 if l.datagram == nil {
515 return
516 }
517
518 for {
519 select {
520 case <-ctx.Done():
521 l.datagram.Clear("lease stopped")
522 return
523 default:
524 }
525
526 conn, err := l.openQUICBackhaulSession(ctx)
527 if err != nil {
528 log.Info().
529 Err(err).
530 Str("component", "sdk-quic-backhaul").
531 Str("address", l.identity.Address).
532 Msg("quic backhaul unavailable; retrying")
533 if !utils.SleepOrDone(ctx, 2*time.Second) {
534 l.datagram.Clear("lease stopped")
535 return
536 }
537 continue
538 }
539
540 log.Info().
541 Str("component", "sdk-quic-backhaul").
542 Str("address", l.identity.Address).
543 Str("remote_addr", conn.RemoteAddr().String()).
544 Msg("quic backhaul connected")
545
546 recvDone, err := l.datagram.BindBackhaul(conn)
547 if err != nil {
548 if ctx.Err() != nil {
549 return
550 }
551 log.Info().
552 Err(err).
553 Str("component", "sdk-quic-backhaul").
554 Str("address", l.identity.Address).
555 Msg("quic backhaul did not bind cleanly; retrying")
556 if !utils.SleepOrDone(ctx, time.Second) {
557 return
558 }
559 continue
560 }
561
562 select {
563 case <-ctx.Done():
564 l.datagram.Clear("lease stopped")
565 return
566 case <-recvDone:
567 }
568
569 if !utils.SleepOrDone(ctx, time.Second) {
570 return
571 }
572 }
573 }
574
575 func (l *listener) openReverseSession(ctx context.Context) (net.Conn, error) {
576 lease, ok := l.leaseSnapshot()
577 if !ok || lease.accessToken == "" {
578 return nil, errors.New("access token is not available")
579 }
580 if l.tlsConfig == nil {
581 return nil, errors.New("relay tls config is unavailable")
582 }
583
584 dialer := &tls.Dialer{
585 NetDialer: &net.Dialer{Timeout: l.dialTimeout},
586 Config: l.tlsConfig.Clone(),
587 }
588
589 conn, err := dialer.DialContext(ctx, "tcp", utils.EnsurePort(l.relayURL.Host))
590 if err != nil {
591 return nil, err
592 }
593
594 req := &http.Request{
595 Method: http.MethodGet,
596 URL: utils.ResolveAPIURL(l.relayURL, types.PathSDKConnect),
597 Host: l.relayURL.Host,
598 Header: make(http.Header),
599 }
600 req.Header.Set(types.HeaderAccessToken, lease.accessToken)
601 req.Header.Set("Connection", "Upgrade")
602 req.Header.Set("Upgrade", "raw")
603
604 if writeErr := req.Write(conn); writeErr != nil {
605 _ = conn.Close()
606 return nil, writeErr
607 }
608
609 reader := bufio.NewReader(conn)
610 resp, err := http.ReadResponse(reader, req)
611 if err != nil {
612 _ = conn.Close()
613 return nil, err
614 }
615 defer resp.Body.Close()
616
617 if resp.StatusCode != http.StatusSwitchingProtocols {
618 apiErr := utils.DecodeAPIRequestError(resp)
619 _ = conn.Close()
620 return nil, apiErr
621 }
622
623 return wrapBufferedConn(conn, reader), nil
624 }
625
626 func (l *listener) openQUICBackhaulSession(ctx context.Context) (*quic.Conn, error) {
627 lease, ok := l.leaseSnapshot()
628 if !ok || lease.accessToken == "" {
629 return nil, errors.New("access token is not available")
630 }
631 if lease.sniPort <= 0 {
632 return nil, errors.New("sni port is not available")
633 }
634 if l.tlsConfig == nil {
635 return nil, errors.New("relay tls config is unavailable")
636 }
637 host := strings.TrimSpace(l.relayURL.Hostname())
638 if host == "" {
639 host = strings.TrimSpace(l.relayURL.Host)
640 }
641 dialAddr := net.JoinHostPort(host, fmt.Sprintf("%d", lease.sniPort))
642 return transport.DialQUICBackhaul(ctx, dialAddr, l.tlsConfig, lease.accessToken)
643 }
644
645 func (l *listener) runRenewLoop(ctx context.Context) error {
646 const wakeThreshold = 10 * time.Second
647
648 for {
649 interval, err := l.renewDelay(time.Now())
650 if err != nil {
651 return err
652 }
653
654 // Round(0) strips the monotonic clock reading so that
655 // time.Since uses wall-clock time. The monotonic clock
656 // freezes during macOS sleep, so without this the elapsed
657 // duration would equal the timer interval, not real time.
658 before := time.Now().Round(0)
659 if interval > 0 {
660 if !utils.SleepOrDone(ctx, interval) {
661 return ctx.Err()
662 }
663 }
664 elapsed := time.Since(before)
665
666 // If the wall-clock jump is much larger than expected, the OS
667 // likely suspended the process (e.g. macOS lid close). The
668 // server-side lease is almost certainly expired, so skip the
669 // normal renew and go straight to re-registration.
670 if elapsed > interval+wakeThreshold {
671 log.Info().
672 Dur("expected", interval).
673 Dur("actual", elapsed).
674 Str("address", l.identity.Address).
675 Msg("system sleep/wake detected; resetting transport and re-registering")
676 return errLeaseRefreshRequired
677 }
678
679 var retries int
680 for {
681 err := l.renewLease(ctx)
682 if err == nil {
683 break
684 }
685 if errors.Is(err, context.Canceled) || errors.Is(err, net.ErrClosed) {
686 return err
687 }
688 if errors.Is(err, errLeaseRefreshRequired) {
689 return err
690 }
691
692 retries++
693 if !l.waitRetry(ctx, "lease renewal", err, retries, 0) {
694 return err
695 }
696 }
697 }
698 }
699
700 func (l *listener) renewDelay(now time.Time) (time.Duration, error) {
701 lease, ok := l.leaseSnapshot()
702 if !ok || lease.accessToken == "" || !now.Before(lease.expiresAt) {
703 return 0, errLeaseRefreshRequired
704 }
705
706 leaseTTL := l.leaseTTL
707 if leaseTTL <= 0 {
708 leaseTTL = defaultLeaseTTL
709 }
710 renewBefore := l.renewBefore
711 if renewBefore <= 0 || renewBefore >= leaseTTL {
712 renewBefore = leaseTTL / 2
713 }
714 if renewBefore <= 0 {
715 renewBefore = time.Second
716 }
717
718 renewAt := lease.expiresAt.Add(-renewBefore)
719 if !now.Before(renewAt) {
720 return 0, nil
721 }
722 return renewAt.Sub(now), nil
723 }
724
725 func (l *listener) renewLease(ctx context.Context) error {
726 lease, ok := l.leaseSnapshot()
727 if !ok || lease.accessToken == "" || !time.Now().Before(lease.expiresAt) {
728 return errLeaseRefreshRequired
729 }
730
731 requestCtx, cancel := context.WithTimeout(ctx, 10*time.Second)
732 defer cancel()
733
734 resp, err := l.renewRegisteredLease(requestCtx, l.leaseTTL, lease.accessToken)
735 if err != nil {
736 if errors.Is(err, &types.APIRequestError{Code: types.APIErrorCodeLeaseNotFound}) {
737 return errLeaseRefreshRequired
738 }
739 return err
740 }
741
742 resp.AccessToken = strings.TrimSpace(resp.AccessToken)
743 if resp.AccessToken == "" {
744 return errors.New("relay did not return renewed access token")
745 }
746 multihopAccessToken := resp.AccessToken
747 var entrySNIPort int
748 if len(lease.hopRoutes) > 0 {
749 multihopAccessToken, entrySNIPort, err = l.registerHopRoutes(requestCtx, resp.ExpiresAt, lease.hopRoutes)
750 if err != nil {
751 return err
752 }
753 }
754 if l.lease == nil {
755 return errLeaseRefreshRequired
756 }
757 _, updated := l.lease.UpdateIf(func(current listenerSnapshot) (listenerSnapshot, bool) {
758 if current.accessToken != lease.accessToken {
759 return current, false
760 }
761 next := current
762 next.accessToken = resp.AccessToken
763 next.expiresAt = resp.ExpiresAt
764 next.multihopAccessToken = multihopAccessToken
765 if entrySNIPort > 0 {
766 next.sniPort = entrySNIPort
767 }
768 return next, true
769 })
770 if !updated {
771 return errLeaseRefreshRequired
772 }
773 return nil
774 }
775
776 func (l *listener) registerAndConfigure(ctx context.Context) error {
777 if err := l.initHTTPTransport(ctx); err != nil {
778 return err
779 }
780
781 resp, hopRoutes, publicHostname, routeHostname, err := l.registerLease(ctx, l.leaseTTL, l.udpEnabled, l.tcpEnabled)
782 if err != nil {
783 return err
784 }
785 resp.AccessToken = strings.TrimSpace(resp.AccessToken)
786 if resp.AccessToken == "" {
787 return errors.New("relay did not return access token")
788 }
789 if l.udpEnabled && !resp.UDPEnabled {
790 _ = l.unregisterLease(context.Background(), resp.AccessToken, hopRoutes)
791 return &types.APIRequestError{
792 Code: types.APIErrorCodeFeatureUnavailable,
793 Message: "relay did not enable required udp support",
794 }
795 }
796 if l.udpEnabled && resp.SNIPort <= 0 {
797 _ = l.unregisterLease(context.Background(), resp.AccessToken, hopRoutes)
798 return errors.New("relay did not return sni port for udp transport")
799 }
800 multihopAccessToken := resp.AccessToken
801 sniPort := resp.SNIPort
802 if len(hopRoutes) > 0 {
803 multihopAccessToken, sniPort, err = l.registerHopRoutes(ctx, resp.ExpiresAt, hopRoutes)
804 if err != nil {
805 _ = l.unregisterLease(context.Background(), resp.AccessToken, hopRoutes)
806 return err
807 }
808 }
809 keylessURL := l.relayURL.String()
810 if multiHop := l.route.MultiHop(); len(multiHop) > 0 {
811 keylessURL = multiHop[0]
812 }
813 publicURLBase := l.relayURL
814 if normalizedKeylessURL, err := utils.NormalizeRelayURL(keylessURL); err == nil {
815 if parsedKeylessURL, parseErr := url.Parse(normalizedKeylessURL); parseErr == nil {
816 publicURLBase = parsedKeylessURL
817 }
818 }
819 echKeys, echConfigList, err := l.tenantECHMaterials(publicHostname, routeHostname)
820 if err != nil {
821 _ = l.unregisterLease(context.Background(), resp.AccessToken, hopRoutes)
822 return err
823 }
824
825 tlsConf, tenantTLSCloser, err := keyless.BuildClientTLSConfig(keylessURL, publicHostname, echKeys, func() http.Header {
826 headers := http.Header{}
827 accessToken := multihopAccessToken
828 if snapshot, ok := l.leaseSnapshot(); ok && snapshot.multihopAccessToken != "" {
829 accessToken = snapshot.multihopAccessToken
830 }
831 headers.Set(types.HeaderAccessToken, accessToken)
832 return headers
833 })
834 if err != nil {
835 _ = l.unregisterLease(context.Background(), resp.AccessToken, hopRoutes)
836 if tenantTLSCloser != nil {
837 _ = tenantTLSCloser.Close()
838 }
839 return err
840 }
841
842 if ctx.Err() != nil {
843 _ = l.unregisterLease(context.Background(), resp.AccessToken, hopRoutes)
844 if tenantTLSCloser != nil {
845 _ = tenantTLSCloser.Close()
846 }
847 return ctx.Err()
848 }
849 next := listenerSnapshot{
850 hostname: publicHostname,
851 echConfigList: echConfigList,
852 udpAddr: resp.UDPAddr,
853 tcpAddr: resp.TCPAddr,
854 accessToken: resp.AccessToken,
855 expiresAt: resp.ExpiresAt,
856 sniPort: sniPort,
857 publicURLBase: publicURLBase,
858 tlsConfig: tlsConf,
859 tlsCloser: tenantTLSCloser,
860 multihopAccessToken: multihopAccessToken,
861 hopRoutes: hopRoutes,
862 }
863 oldLease := l.lease.Swap(next)
864 if oldLease.tlsCloser != nil {
865 _ = oldLease.tlsCloser.Close()
866 }
867 if l.udpEnabled && l.datagram != nil {
868 l.datagram.Clear("lease updated")
869 }
870 relayURL := l.relayURL.String()
871 if l.relaySet != nil && relayURL != "" {
872 l.relaySet.ConfirmRelayURL(relayURL)
873 }
874 if len(echConfigList) > 0 {
875 log.Info().
876 Str("address", l.identity.Address).
877 Str("route_hostname", routeHostname).
878 Str("ech_config_list_base64", base64.StdEncoding.EncodeToString(echConfigList)).
879 Msg("tenant ech config ready")
880 }
881 return nil
882 }
883
884 func (l *listener) tenantECHMaterials(publicHostname, routeHostname string) ([]tls.EncryptedClientHelloKey, []byte, error) {
885 if routeHostname == "" {
886 return nil, nil, nil
887 }
888 echSeed, err := identity.DeriveToken(l.identity, "tenant-ech", publicHostname, routeHostname)
889 if err != nil {
890 return nil, nil, fmt.Errorf("derive tenant ech seed: %w", err)
891 }
892 echKeys, echConfigList, err := keyless.EncryptedClientHelloMaterials(echSeed, routeHostname)
893 if err != nil {
894 return nil, nil, fmt.Errorf("prepare tenant ech materials: %w", err)
895 }
896 return echKeys, echConfigList, nil
897 }
898
899 func (l *listener) waitRetry(ctx context.Context, operation string, err error, retries, reverseSessionSlot int) bool {
900 if ctx.Err() != nil {
901 return false
902 }
903
904 relayURL := ""
905 if l.relayURL != nil {
906 relayURL = l.relayURL.String()
907 }
908 logger := log.With().
909 Str("relay_url", relayURL).
910 Str("operation", operation).
911 Str("address", l.identity.Address).
912 Logger()
913 if reverseSessionSlot > 0 {
914 logger = logger.With().Int("reverse_session_slot", reverseSessionSlot).Logger()
915 }
916
917 if l.retryCount > 0 && retries > l.retryCount {
918 if l.relaySet != nil && relayURL != "" {
919 l.relaySet.UnconfirmRelayURL(relayURL)
920 l.relaySet.RecordActiveFailure(relayURL, 1)
921 l.relaySet.DropRelayURLFromActivePool(relayURL)
922 }
923 logger.Error().
924 Err(err).
925 Int("retry_count", l.retryCount).
926 Msg("retry budget exhausted")
927 return false
928 }
929
930 logger.Debug().
931 Err(err).
932 Int("retry_attempt", retries).
933 Int("retry_count", l.retryCount).
934 Dur("retry_wait", l.retryWait).
935 Msg("operation failed; retrying")
936
937 return utils.SleepOrDone(ctx, l.retryWait)
938 }
939
940 type bufferedConn struct {
941 net.Conn
942 reader *bytes.Reader
943 }
944
945 func wrapBufferedConn(conn net.Conn, reader *bufio.Reader) net.Conn {
946 if reader == nil || reader.Buffered() == 0 {
947 return conn
948 }
949 buf := make([]byte, reader.Buffered())
950 if _, err := io.ReadFull(reader, buf); err != nil {
951 return conn
952 }
953 return &bufferedConn{Conn: conn, reader: bytes.NewReader(buf)}
954 }
955
956 func (c *bufferedConn) Read(p []byte) (int, error) {
957 if c.reader != nil && c.reader.Len() > 0 {
958 return c.reader.Read(p)
959 }
960 return c.Conn.Read(p)
961 }