main
go 855 lines 23.5 KB
Raw
1 package portal
2
3 import (
4 "context"
5 "crypto/tls"
6 "errors"
7 "fmt"
8 "io"
9 "net"
10 "net/http"
11 "net/http/pprof"
12 "strings"
13 "sync"
14 "time"
15
16 "github.com/gosuda/keyless_tls/relay/l4"
17 "github.com/quic-go/quic-go"
18 "github.com/rs/zerolog/log"
19 "golang.org/x/sync/errgroup"
20
21 "github.com/gosuda/portal-tunnel/v2/portal/acme"
22 "github.com/gosuda/portal-tunnel/v2/portal/auth"
23 "github.com/gosuda/portal-tunnel/v2/portal/discovery"
24 "github.com/gosuda/portal-tunnel/v2/portal/identity"
25 "github.com/gosuda/portal-tunnel/v2/portal/keyless"
26 "github.com/gosuda/portal-tunnel/v2/portal/overlay"
27 "github.com/gosuda/portal-tunnel/v2/portal/policy"
28 "github.com/gosuda/portal-tunnel/v2/portal/transport"
29 "github.com/gosuda/portal-tunnel/v2/types"
30 "github.com/gosuda/portal-tunnel/v2/utils"
31 )
32
33 const (
34 defaultClaimTimeout = 10 * time.Second
35 defaultClientHelloWait = 2 * time.Second
36 defaultControlBodyLimit = 4 << 20
37 defaultHopOpenRetryWait = 250 * time.Millisecond
38 DefaultPProfListenAddr = "127.0.0.1:6060"
39 )
40
41 type ServerConfig struct {
42 PortalURL string
43 IdentityPath string
44 Bootstraps []string
45 DiscoveryEnabled bool
46 WireGuardPort int
47 APIPort int
48 SNIPort int
49 APIListenAddr string
50 SNIListenAddr string
51 TrustProxyHeaders bool
52 TrustedProxyCIDRs string
53 UDPEnabled bool
54 TCPEnabled bool
55 MinPort int
56 MaxPort int
57 PProfEnabled bool
58 PProfListenAddr string
59 X402Enabled bool
60 X402Testnet bool
61 X402PayTo string
62 ACME acme.Config
63 }
64
65 func normalizeServerConfig(cfg ServerConfig) (ServerConfig, error) {
66 cfg.PortalURL = strings.TrimSuffix(strings.TrimSpace(cfg.PortalURL), "/")
67 cfg.IdentityPath = identity.ResolveRelayStateDir(cfg.IdentityPath)
68 if cfg.IdentityPath == "" {
69 return ServerConfig{}, errors.New("identity path is required")
70 }
71
72 selfRelayURL, err := utils.NormalizeRelayURL(cfg.PortalURL)
73 if err != nil {
74 return ServerConfig{}, fmt.Errorf("normalize portal url: %w", err)
75 }
76 if utils.PortalRootHost(selfRelayURL) == "" {
77 return ServerConfig{}, errors.New("root host is required")
78 }
79
80 bootstraps, err := utils.NormalizeRelayURLs(cfg.Bootstraps...)
81 if err != nil {
82 return ServerConfig{}, fmt.Errorf("normalize bootstraps: %w", err)
83 }
84 cfg.PortalURL = selfRelayURL
85 cfg.Bootstraps = bootstraps
86 cfg.Bootstraps = utils.RemoveRelayURL(cfg.Bootstraps, selfRelayURL)
87
88 cfg.APIPort = utils.IntOrDefault(cfg.APIPort, 4017)
89 cfg.SNIPort = utils.IntOrDefault(cfg.SNIPort, 443)
90 cfg.WireGuardPort = utils.IntOrDefault(cfg.WireGuardPort, overlay.DefaultListenPort)
91 cfg.APIListenAddr = utils.StringOrDefault(cfg.APIListenAddr, fmt.Sprintf(":%d", cfg.APIPort))
92 cfg.SNIListenAddr = utils.StringOrDefault(cfg.SNIListenAddr, fmt.Sprintf(":%d", cfg.SNIPort))
93 if cfg.PProfEnabled {
94 cfg.PProfListenAddr = utils.StringOrDefault(strings.TrimSpace(cfg.PProfListenAddr), DefaultPProfListenAddr)
95 }
96 cfg.X402PayTo = strings.TrimSpace(cfg.X402PayTo)
97 hasPortRange := cfg.MinPort > 0 && cfg.MaxPort > 0
98 if cfg.UDPEnabled || cfg.TCPEnabled {
99 switch {
100 case !hasPortRange:
101 return ServerConfig{}, errors.New("udp and tcp relay transport require a valid min port and max port range")
102 case cfg.MinPort > 65535 || cfg.MaxPort > 65535:
103 return ServerConfig{}, errors.New("min port and max port must be between 1 and 65535")
104 case cfg.MinPort > cfg.MaxPort:
105 return ServerConfig{}, errors.New("min port must be less than or equal to max port")
106 }
107 }
108
109 cfg.UDPEnabled = cfg.UDPEnabled && cfg.hasLeasePortRange()
110 cfg.TCPEnabled = cfg.TCPEnabled && cfg.hasLeasePortRange()
111 return cfg, nil
112 }
113
114 func (cfg ServerConfig) snapshot() ServerConfig {
115 cfg.Bootstraps = utils.CloneSlice(cfg.Bootstraps)
116 return cfg
117 }
118
119 func (cfg ServerConfig) hasLeasePortRange() bool {
120 return cfg.MinPort > 0 && cfg.MaxPort > 0 && cfg.MinPort <= 65535 && cfg.MaxPort <= 65535 && cfg.MinPort <= cfg.MaxPort
121 }
122
123 type Server struct {
124 cancel context.CancelFunc
125 group *errgroup.Group
126 shutdownOnce sync.Once
127
128 cfg *utils.Snapshot[ServerConfig]
129 identity types.RelayIdentity
130 authority identity.Authority
131 acmeManager *acme.Manager
132 proxy proxy
133
134 apiListener net.Listener
135 sniListener net.Listener
136 apiServer *http.Server
137 apiTLSClose io.Closer
138 pprofListener net.Listener
139 pprofServer *http.Server
140 quicBackhaul *quic.Listener
141
142 overlay *overlay.Overlay
143 relaySet *discovery.RelaySet
144 announceLimiter *discovery.AnnounceLimiter
145 registry *leaseRegistry
146 }
147
148 func NewServer(cfg ServerConfig) (*Server, error) {
149 cfg, err := normalizeServerConfig(cfg)
150 if err != nil {
151 return nil, err
152 }
153
154 relayIdentity, err := identity.LoadOrCreateRelayIdentity(cfg.IdentityPath, utils.PortalRootHost(cfg.PortalURL), cfg.DiscoveryEnabled)
155 if err != nil {
156 return nil, fmt.Errorf("load relay identity: %w", err)
157 }
158 relayAuthority, err := identity.NewLocalAuthority(relayIdentity.Identity)
159 if err != nil {
160 return nil, fmt.Errorf("load relay authority: %w", err)
161 }
162 registry, err := newLeaseRegistry(cfg.UDPEnabled, cfg.TCPEnabled, cfg.MinPort, cfg.MaxPort, relayIdentity.Name, cfg.SNIPort, relayAuthority, cfg.PortalURL, cfg.TrustProxyHeaders, cfg.TrustedProxyCIDRs)
163 if err != nil {
164 return nil, err
165 }
166 var relaySet *discovery.RelaySet
167 if cfg.DiscoveryEnabled {
168 cfg.Bootstraps, err = utils.ResolvePortalRelayURLs(cfg.Bootstraps, true)
169 if err != nil {
170 return nil, fmt.Errorf("resolve discovery bootstraps: %w", err)
171 }
172 cfg.Bootstraps = utils.RemoveRelayURL(cfg.Bootstraps, cfg.PortalURL)
173 relaySet = discovery.NewRelaySet(cfg.Bootstraps)
174 }
175
176 server := &Server{
177 cfg: utils.NewSnapshot(cfg, ServerConfig.snapshot),
178 identity: relayIdentity,
179 authority: relayAuthority,
180 registry: registry,
181 relaySet: relaySet,
182 announceLimiter: discovery.NewAnnounceLimiter(0, 0),
183 }
184 server.registry.proxy = &server.proxy
185 return server, nil
186 }
187
188 func (s *Server) config() ServerConfig {
189 return s.cfg.Load()
190 }
191
192 func (s *Server) SetUDPPolicy(enabled bool, maxLeases int) {
193 if enabled && !s.config().hasLeasePortRange() {
194 enabled = false
195 }
196 if runtime := s.PolicyRuntime(); runtime != nil {
197 runtime.SetUDPPolicy(enabled, maxLeases)
198 }
199 s.cfg.UpdateCopy(func(cfg *ServerConfig) {
200 cfg.UDPEnabled = enabled
201 })
202 }
203
204 func (s *Server) SetTCPPortPolicy(enabled bool, maxLeases int) {
205 if enabled && !s.config().hasLeasePortRange() {
206 enabled = false
207 }
208 if runtime := s.PolicyRuntime(); runtime != nil {
209 runtime.SetTCPPortPolicy(enabled, maxLeases)
210 }
211 s.cfg.UpdateCopy(func(cfg *ServerConfig) {
212 cfg.TCPEnabled = enabled
213 })
214 }
215
216 func (s *Server) supportsUDP() bool {
217 runtime := s.PolicyRuntime()
218 if runtime == nil || !runtime.IsUDPEnabled() {
219 return false
220 }
221 return s.group == nil || s.quicBackhaul != nil
222 }
223
224 func (s *Server) supportsTCP() bool {
225 runtime := s.PolicyRuntime()
226 return runtime != nil && runtime.IsTCPPortEnabled()
227 }
228
229 func (s *Server) Start(ctx context.Context, apiMux *http.ServeMux) error {
230 if s.group != nil {
231 return errors.New("server already started")
232 }
233 cfg := s.config()
234 apiTLS, acmeManager, err := s.prepareAPITLS(ctx)
235 if err != nil {
236 return err
237 }
238
239 serverCtx, cancel := context.WithCancel(ctx)
240 started := false
241 var apiListener net.Listener
242 var sniListener net.Listener
243 var apiServer *http.Server
244 var apiCloser io.Closer
245 var pprofListener net.Listener
246 var pprofServer *http.Server
247 var ov *overlay.Overlay
248 var quicBackhaul *quic.Listener
249 defer func() {
250 if started {
251 return
252 }
253 acmeManager.Stop()
254 if ov != nil {
255 _ = ov.Shutdown(context.Background())
256 }
257 if apiServer != nil {
258 _ = apiServer.Close()
259 }
260 if pprofServer != nil {
261 _ = pprofServer.Close()
262 }
263 if pprofListener != nil {
264 _ = pprofListener.Close()
265 }
266 if apiCloser != nil {
267 _ = apiCloser.Close()
268 }
269 if sniListener != nil {
270 _ = sniListener.Close()
271 }
272 if apiListener != nil {
273 _ = apiListener.Close()
274 }
275 cancel()
276 }()
277 var listenConfig net.ListenConfig
278
279 apiListener, err = listenConfig.Listen(serverCtx, "tcp", cfg.APIListenAddr)
280 if err != nil {
281 return fmt.Errorf("listen api: %w", err)
282 }
283 sniListener, err = listenConfig.Listen(serverCtx, "tcp", cfg.SNIListenAddr)
284 if err != nil {
285 return fmt.Errorf("listen sni: %w", err)
286 }
287
288 group, groupCtx := errgroup.WithContext(serverCtx)
289 wrappedAPIListener, apiServer, apiCloser, err := s.newAPIServer(apiListener, apiMux, apiTLS)
290 if err != nil {
291 return err
292 }
293 if cfg.PProfEnabled {
294 pprofListener, err = listenConfig.Listen(serverCtx, "tcp", cfg.PProfListenAddr)
295 if err != nil {
296 return fmt.Errorf("listen pprof: %w", err)
297 }
298 pprofMux := http.NewServeMux()
299 pprofMux.HandleFunc("/debug/pprof/", pprof.Index)
300 pprofMux.HandleFunc("/debug/pprof/cmdline", pprof.Cmdline)
301 pprofMux.HandleFunc("/debug/pprof/profile", pprof.Profile)
302 pprofMux.HandleFunc("/debug/pprof/symbol", pprof.Symbol)
303 pprofMux.HandleFunc("/debug/pprof/trace", pprof.Trace)
304 pprofServer = &http.Server{
305 Handler: pprofMux,
306 ReadHeaderTimeout: 10 * time.Second,
307 }
308 }
309
310 if s.relaySet != nil && strings.TrimSpace(s.identity.WireGuardPrivateKey) != "" {
311 ov, err = s.startOverlay()
312 if err != nil {
313 return err
314 }
315 }
316 if cfg.UDPEnabled {
317 quicBackhaul, err = s.newQUICBackhaulListener(apiTLS)
318 if err != nil {
319 log.Warn().Err(err).Msg("quic backhaul listener disabled")
320 quicBackhaul = nil
321 }
322 }
323
324 s.apiListener = wrappedAPIListener
325 s.sniListener = sniListener
326 s.apiServer = apiServer
327 s.apiTLSClose = apiCloser
328 s.pprofListener = pprofListener
329 s.pprofServer = pprofServer
330 s.acmeManager = acmeManager
331 s.cancel = cancel
332 s.group = group
333 s.overlay = ov
334 s.quicBackhaul = quicBackhaul
335 started = true
336
337 group.Go(s.runAPIServer)
338 if s.pprofServer != nil {
339 group.Go(s.runPProfServer)
340 }
341 group.Go(func() error { return s.runPublicIngress(groupCtx) })
342 if s.overlay != nil {
343 group.Go(func() error { return s.overlay.Serve(groupCtx) })
344 }
345 if s.quicBackhaul != nil {
346 group.Go(s.runQUICBackhaulListener)
347 }
348 group.Go(func() error { return s.runRegistryJanitor(groupCtx, 5*time.Second) })
349 if cfg.DiscoveryEnabled {
350 group.Go(func() error { return s.runRelayDiscoveryLoop(groupCtx) })
351 }
352 s.acmeManager.Start(serverCtx)
353 group.Go(func() error {
354 <-groupCtx.Done()
355 shutdownCtx, cancel := context.WithTimeout(context.Background(), 5*time.Second)
356 defer cancel()
357 return s.Shutdown(shutdownCtx)
358 })
359
360 logEvent := log.Info().
361 Str("api_addr", utils.HostPortOrLoopback(s.apiListener.Addr().String())).
362 Str("sni_addr", s.sniListener.Addr().String()).
363 Str("root_host", s.identity.Name).
364 Str("acme_dns_provider", cfg.ACME.DNSProvider).
365 Int("min_port", cfg.MinPort).
366 Int("max_port", cfg.MaxPort).
367 Bool("discovery_enabled", cfg.DiscoveryEnabled).
368 Bool("wireguard_enabled", s.overlay != nil).
369 Bool("multihop_enabled", s.overlay != nil).
370 Bool("udp_enabled", s.quicBackhaul != nil).
371 Bool("tcp_enabled", s.supportsTCP()).
372 Bool("api_ech_enabled", len(apiTLS.EncryptedClientHelloKeys) > 0).
373 Bool("pprof_enabled", s.pprofServer != nil)
374 if s.pprofListener != nil {
375 logEvent = logEvent.Str("pprof_addr", utils.HostPortOrLoopback(s.pprofListener.Addr().String()))
376 }
377 if s.quicBackhaul != nil {
378 logEvent = logEvent.Str("internal_quic_backhaul_addr", s.quicBackhaul.Addr().String())
379 }
380 logEvent.Msg("relay server started")
381
382 return nil
383 }
384
385 func (s *Server) Wait() error {
386 if s.group == nil {
387 return nil
388 }
389 err := s.group.Wait()
390 if errors.Is(err, context.Canceled) {
391 return nil
392 }
393 return err
394 }
395
396 func (s *Server) PolicyRuntime() *policy.Runtime {
397 if s == nil || s.registry == nil {
398 return nil
399 }
400 return s.registry.policy
401 }
402
403 func (s *Server) PortalURL() string {
404 if s == nil {
405 return ""
406 }
407 return s.config().PortalURL
408 }
409
410 func (s *Server) PublicLeases() []types.Lease {
411 if s == nil || s.registry == nil {
412 return nil
413 }
414 return s.registry.PublicLeases(time.Now())
415 }
416
417 func (s *Server) PolicyLeases() []types.PolicyLease {
418 if s == nil || s.registry == nil {
419 return nil
420 }
421 return s.registry.PolicyLeases(time.Now())
422 }
423
424 func (s *Server) RelayIdentity() types.RelayIdentity {
425 if s == nil {
426 return types.RelayIdentity{}
427 }
428 return s.identity.Copy()
429 }
430
431 func (s *Server) Shutdown(ctx context.Context) error {
432 var shutdownErr error
433 s.shutdownOnce.Do(func() {
434 if s.cancel != nil {
435 s.cancel()
436 }
437
438 records := s.registry.CloseAll()
439 for _, record := range records {
440 record.deleteDNS(ctx, s.acmeManager, true)
441 }
442
443 if s.quicBackhaul != nil {
444 _ = s.quicBackhaul.Close()
445 }
446 if s.sniListener != nil {
447 if err := s.sniListener.Close(); err != nil && !errors.Is(err, net.ErrClosed) {
448 shutdownErr = err
449 }
450 }
451 if s.apiServer != nil {
452 if err := s.apiServer.Shutdown(ctx); err != nil && shutdownErr == nil {
453 shutdownErr = err
454 }
455 }
456 if s.pprofServer != nil {
457 if err := s.pprofServer.Shutdown(ctx); err != nil && shutdownErr == nil {
458 shutdownErr = err
459 }
460 }
461 if s.overlay != nil {
462 if err := s.overlay.Shutdown(ctx); err != nil && shutdownErr == nil {
463 shutdownErr = err
464 }
465 }
466 if s.apiTLSClose != nil {
467 _ = s.apiTLSClose.Close()
468 }
469 if s.acmeManager != nil {
470 s.acmeManager.Stop()
471 }
472 })
473 return shutdownErr
474 }
475
476 func (s *Server) prepareAPITLS(ctx context.Context) (keyless.TLSMaterialConfig, *acme.Manager, error) {
477 cfg := s.config()
478 acmeCfg := cfg.ACME
479 if baseDomain := utils.NormalizeHostname(acmeCfg.BaseDomain); baseDomain != "" && baseDomain != s.identity.Name {
480 return keyless.TLSMaterialConfig{}, nil, fmt.Errorf("acme base domain %q does not match portal root host %q", acmeCfg.BaseDomain, s.identity.Name)
481 }
482 acmeCfg.BaseDomain = s.identity.Name
483 if strings.TrimSpace(acmeCfg.ENSGaslessAddress) == "" {
484 acmeCfg.ENSGaslessAddress = s.identity.Address
485 }
486
487 manager, err := acme.NewManager(acmeCfg)
488 if err != nil {
489 return keyless.TLSMaterialConfig{}, nil, fmt.Errorf("create acme manager: %w", err)
490 }
491
492 certPEM, keyPEM, err := manager.EnsureTLSMaterial(ctx)
493 if err != nil {
494 manager.Stop()
495 return keyless.TLSMaterialConfig{}, nil, fmt.Errorf("ensure relay certificate: %w", err)
496 }
497
498 apiTLS := keyless.TLSMaterialConfig{
499 CertPEM: certPEM,
500 KeyPEM: keyPEM,
501 }
502 echSeed, err := identity.DeriveToken(
503 s.identity.Identity,
504 "relay-ech",
505 s.identity.EncryptedClientHelloSeed,
506 s.identity.Name,
507 )
508 if err != nil {
509 manager.Stop()
510 return keyless.TLSMaterialConfig{}, nil, fmt.Errorf("derive relay ech seed: %w", err)
511 }
512 echKeys, echConfigList, err := keyless.EncryptedClientHelloMaterials(echSeed, s.identity.Name)
513 if err != nil {
514 manager.Stop()
515 return keyless.TLSMaterialConfig{}, nil, fmt.Errorf("prepare ech materials: %w", err)
516 }
517 if len(echKeys) > 0 {
518 apiTLS.EncryptedClientHelloKeys = echKeys
519 if err := manager.SyncECHConfig(ctx, s.identity.Name, echConfigList, cfg.SNIPort); err != nil {
520 log.Warn().
521 Err(err).
522 Str("hostname", s.identity.Name).
523 Msg("publish relay ech dns record")
524 }
525 }
526
527 return apiTLS, manager, nil
528 }
529
530 func (s *Server) runAPIServer() error {
531 err := s.apiServer.Serve(s.apiListener)
532 if err == nil || errors.Is(err, http.ErrServerClosed) || errors.Is(err, net.ErrClosed) {
533 return nil
534 }
535 return err
536 }
537
538 func (s *Server) runPProfServer() error {
539 err := s.pprofServer.Serve(s.pprofListener)
540 if err == nil || errors.Is(err, http.ErrServerClosed) || errors.Is(err, net.ErrClosed) {
541 return nil
542 }
543 return err
544 }
545
546 func (s *Server) runPublicIngress(ctx context.Context) error {
547 for {
548 conn, err := s.sniListener.Accept()
549 switch {
550 case err == nil:
551 go func(conn net.Conn) {
552 clientHello, wrappedConn, err := l4.InspectClientHello(conn, defaultClientHelloWait)
553 if err != nil {
554 if wrappedConn != nil {
555 _ = wrappedConn.Close()
556 } else {
557 _ = conn.Close()
558 }
559 return
560 }
561
562 serverName := utils.NormalizeHostname(clientHello.ServerName)
563 if serverName == "" {
564 _ = wrappedConn.Close()
565 return
566 }
567
568 if serverName == s.identity.Name {
569 if s.apiListener == nil {
570 _ = wrappedConn.Close()
571 return
572 }
573 dialer := &net.Dialer{Timeout: 5 * time.Second}
574 upstream, err := dialer.DialContext(ctx, "tcp", utils.HostPortOrLoopback(s.apiListener.Addr().String()))
575 if err != nil {
576 _ = wrappedConn.Close()
577 return
578 }
579 s.proxy.bridge(wrappedConn, upstream, "", nil)
580 return
581 }
582
583 record, ok := s.registry.Lookup(serverName)
584 if !ok {
585 _ = wrappedConn.Close()
586 return
587 }
588 if err := s.bridgeLeaseConn(ctx, wrappedConn, record); err != nil {
589 log.Warn().Err(err).Msg("bridge public ingress")
590 _ = wrappedConn.Close()
591 return
592 }
593 }(conn)
594 case errors.Is(err, net.ErrClosed):
595 return nil
596 default:
597 if ctxErr := ctx.Err(); ctxErr != nil {
598 return ctxErr
599 }
600 return fmt.Errorf("accept sni connection: %w", err)
601 }
602 }
603 }
604
605 func (s *Server) bridgeLeaseConn(ctx context.Context, conn net.Conn, record *leaseRecord) error {
606 if record.isExpired(time.Now()) {
607 return errLeaseNotFound
608 }
609 if overlayIPv4, forwardToken, hasNextHop := record.nextHop(); hasNextHop {
610 switch {
611 case s.overlay == nil:
612 return errors.New("relay overlay is unavailable")
613 case overlayIPv4 == "":
614 return errors.New("next hop overlay ipv4 is required")
615 case forwardToken == "":
616 return errors.New("next hop token is required")
617 }
618
619 openCtx, cancel := context.WithTimeout(ctx, defaultClaimTimeout)
620 defer cancel()
621 var next net.Conn
622 var lastErr error
623 for {
624 var err error
625 next, err = s.overlay.OpenHopStream(openCtx, overlayIPv4, forwardToken)
626 if err == nil {
627 break
628 }
629 lastErr = err
630 if errors.Is(err, net.ErrClosed) {
631 return fmt.Errorf("open next hop stream: %w", err)
632 }
633 if !utils.SleepOrDone(openCtx, defaultHopOpenRetryWait) {
634 return fmt.Errorf("open next hop stream within %s: %w", defaultClaimTimeout, errors.Join(lastErr, openCtx.Err()))
635 }
636 }
637 s.proxy.bridge(conn, next, "", nil)
638 return nil
639 }
640 if record.stream == nil {
641 return errors.New("lease stream is not ready")
642 }
643 if !s.registry.policy.IsIdentityRoutable(record.Key()) {
644 return errLeaseRejected
645 }
646 claimCtx, cancel := context.WithTimeout(ctx, defaultClaimTimeout)
647 session, err := record.stream.Claim(claimCtx)
648 cancel()
649 if err != nil {
650 return fmt.Errorf("claim lease stream: %w", err)
651 }
652 s.proxy.bridge(conn, session, record.Key(), s.registry.policy.BPSManager())
653 return nil
654 }
655
656 func (s *Server) runRegistryJanitor(ctx context.Context, interval time.Duration) error {
657 if interval <= 0 {
658 return errors.New("janitor interval must be positive")
659 }
660
661 ticker := time.NewTicker(interval)
662 defer ticker.Stop()
663
664 for {
665 select {
666 case <-ctx.Done():
667 return nil
668 case <-ticker.C:
669 records := s.registry.cleanupExpired(time.Now())
670 for _, record := range records {
671 record.deleteDNS(ctx, s.acmeManager, true)
672 }
673 }
674 }
675 }
676
677 func (s *Server) newQUICBackhaulListener(apiTLS keyless.TLSMaterialConfig) (*quic.Listener, error) {
678 if len(apiTLS.KeyPEM) == 0 {
679 return nil, fmt.Errorf("quic backhaul requires api tls key")
680 }
681 tlsCert, err := tls.X509KeyPair(apiTLS.CertPEM, apiTLS.KeyPEM)
682 if err != nil {
683 return nil, fmt.Errorf("parse quic backhaul tls keypair: %w", err)
684 }
685 return transport.ListenQUICBackhaul(s.config().SNIListenAddr, tlsCert)
686 }
687
688 func (s *Server) runQUICBackhaulListener() error {
689 if s.quicBackhaul == nil {
690 return nil
691 }
692 for {
693 conn, err := s.quicBackhaul.Accept(context.Background())
694 if err != nil {
695 if errors.Is(err, quic.ErrServerClosed) || errors.Is(err, net.ErrClosed) {
696 return nil
697 }
698 return err
699 }
700 go s.handleQUICBackhaulConn(conn)
701 }
702 }
703
704 func (s *Server) handleQUICBackhaulConn(conn *quic.Conn) {
705 control, err := transport.AcceptQUICBackhaulControl(context.Background(), conn)
706 if err != nil {
707 _ = conn.CloseWithError(1, "control read failed")
708 return
709 }
710
711 lease, err := s.registry.admitLeaseByToken(control.AccessToken, true)
712 if err != nil {
713 code, reason := types.APIErrorCodeInvalidRequest, "invalid control message"
714 switch {
715 case errors.Is(err, errLeaseNotFound):
716 code, reason = types.APIErrorCodeLeaseNotFound, "lease not found"
717 case errors.Is(err, errLeaseRejected):
718 code, reason = types.APIErrorCodeLeaseRejected, "lease rejected"
719 case errors.Is(err, errUnauthorized):
720 code, reason = types.APIErrorCodeUnauthorized, "unauthorized"
721 case errors.Is(err, errTransportMismatch):
722 code, reason = types.APIErrorCodeTransportMismatch, "transport mismatch"
723 }
724 _ = control.Reject(code, reason)
725 return
726 }
727
728 if err := lease.datagram.BindBackhaul(conn); err != nil {
729 _ = control.Reject("broker_closed", "broker closed")
730 return
731 }
732
733 _ = control.Accept()
734 s.registry.Touch(lease.Key(), conn.RemoteAddr().String(), time.Now())
735 log.Info().
736 Str("component", "quic-backhaul-listener").
737 Str("address", lease.Address).
738 Str("lease_name", lease.Name).
739 Str("remote_addr", conn.RemoteAddr().String()).
740 Msg("quic backhaul connected")
741 }
742
743 func (s *Server) startOverlay() (*overlay.Overlay, error) {
744 cfg := s.config()
745 peerMux := http.NewServeMux()
746 peerMux.HandleFunc(types.PathRoot, s.handleRoot)
747 peerMux.HandleFunc(types.PathHealthz, s.handleHealthz)
748 if cfg.DiscoveryEnabled {
749 peerMux.HandleFunc(types.PathDiscovery, s.handleRelayDiscovery)
750 }
751
752 ov, err := overlay.NewOverlay(overlay.Config{
753 PrivateKey: s.identity.WireGuardPrivateKey,
754 PublicKey: s.identity.WireGuardPublicKey,
755 ListenPort: cfg.WireGuardPort,
756 }, peerMux, nil)
757 if err != nil {
758 return nil, fmt.Errorf("start wireguard overlay: %w", err)
759 }
760
761 ov.SetStreamHandler(func(ctx context.Context, stream overlay.HopStream) {
762 s.registry.mu.RLock()
763 record := s.registry.recordByHopToken(stream.Token, time.Now())
764 s.registry.mu.RUnlock()
765 if record == nil {
766 log.Warn().Str("remote_addr", stream.RemoteAddr).Msg("hop stream rejected")
767 _ = stream.Conn.Close()
768 return
769 }
770 hopRole := "exit"
771 if record.isHopMiddle() {
772 hopRole = "middle"
773 }
774 log.Info().Str("remote_addr", stream.RemoteAddr).Str("hop_role", hopRole).Msg("hop stream received")
775
776 if err := s.bridgeLeaseConn(ctx, stream.Conn, record); err != nil {
777 log.Warn().Err(err).Str("remote_addr", stream.RemoteAddr).Msg("hop stream bridge failed")
778 _ = stream.Conn.Close()
779 }
780 })
781
782 if err := ov.Sync(s.relaySet.OverlayPeerDescriptor()); err != nil {
783 _ = ov.Shutdown(context.Background())
784 return nil, fmt.Errorf("sync wireguard peers: %w", err)
785 }
786
787 return ov, nil
788 }
789
790 func (s *Server) runRelayDiscoveryLoop(ctx context.Context) error {
791 if s.relaySet == nil {
792 <-ctx.Done()
793 return nil
794 }
795 refresher := discovery.NewRefresher(s.relaySet, s.overlay)
796 ticker := time.NewTicker(discovery.DiscoveryPollInterval)
797 defer ticker.Stop()
798
799 for {
800 now := time.Now().UTC()
801 self, err := s.newSelfDescriptor(now)
802 if err != nil {
803 return fmt.Errorf("build relay discovery descriptor: %w", err)
804 }
805 if err := refresher.Refresh(ctx, &self); err != nil {
806 if ctx.Err() != nil {
807 return nil
808 }
809 return err
810 }
811 if ctx.Err() != nil {
812 return nil
813 }
814
815 select {
816 case <-ctx.Done():
817 return nil
818 case <-ticker.C:
819 }
820 }
821 }
822
823 func (s *Server) newSelfDescriptor(now time.Time) (types.RelayDescriptor, error) {
824 if now.IsZero() {
825 now = time.Now().UTC()
826 } else {
827 now = now.UTC()
828 }
829 cfg := s.config()
830
831 var wireGuardPublicKey string
832 var wireGuardPort int
833 supportsOverlay := false
834 if s.overlay != nil {
835 cfg := s.overlay.Config()
836 wireGuardPublicKey = cfg.PublicKey
837 wireGuardPort = cfg.ListenPort
838 supportsOverlay = true
839 }
840
841 return auth.SignRelayDescriptor(types.RelayDescriptor{
842 Address: s.identity.Address,
843 Version: types.DiscoveryVersion,
844 IssuedAt: now,
845 ExpiresAt: now.Add(discovery.DiscoveryDescriptorTTL),
846 APIHTTPSAddr: cfg.PortalURL,
847 WireGuardPublicKey: wireGuardPublicKey,
848 WireGuardPort: wireGuardPort,
849 SupportsOverlay: supportsOverlay,
850 SupportsUDP: s.supportsUDP(),
851 SupportsTCP: s.supportsTCP(),
852 ActiveConnections: s.proxy.activeConnectionCount(),
853 TCPBPS: s.proxy.currentTCPBPS(now),
854 }, s.authority)
855 }