Refactor lease registry and server for improved lease management

Kim committed Apr 29, 2026 at 15:06 UTC d162473db80ebcf91c23bf601fbaa99b42eeba37
5 files changed +471 -515
portal/api_server.go
+34 -315
@@ -12,13 +12,10 @@ import (
12 "strings"
13 "time"
14
15 - "github.com/quic-go/quic-go"
15 "github.com/rs/zerolog/log"
16
17 "github.com/gosuda/portal-tunnel/v2/portal/auth"
19 - "github.com/gosuda/portal-tunnel/v2/portal/discovery"
18 "github.com/gosuda/portal-tunnel/v2/portal/keyless"
21 - "github.com/gosuda/portal-tunnel/v2/portal/transport"
19 "github.com/gosuda/portal-tunnel/v2/types"
20 "github.com/gosuda/portal-tunnel/v2/utils"
21 )
@@ -41,14 +38,14 @@ var (
38 errUnauthorized = &apiError{types.APIErrorCodeUnauthorized, "unauthorized", http.StatusForbidden}
39 errUDPDisabled = &apiError{types.APIErrorCodeUDPDisabled, "udp disabled", http.StatusForbidden}
40 errUDPCapacityExceeded = &apiError{types.APIErrorCodeUDPCapacityExceeded, "udp capacity exceeded", http.StatusServiceUnavailable}
41 + errUDPPortExhausted = &apiError{types.APIErrorCodeUDPPortExhausted, "no udp ports available", http.StatusServiceUnavailable}
42 errTCPPortDisabled = &apiError{types.APIErrorCodeTCPPortDisabled, "tcp port disabled", http.StatusForbidden}
43 errTCPPortCapacityExceeded = &apiError{types.APIErrorCodeTCPPortCapacityExceeded, "tcp port capacity exceeded", http.StatusServiceUnavailable}
44 errTCPPortExhausted = &apiError{types.APIErrorCodeTCPPortExhausted, "no tcp ports available", http.StatusServiceUnavailable}
45 )
46
47 func writeAPIErrorResponse(w http.ResponseWriter, err error) {
50 - var ae *apiError
51 - if errors.As(err, &ae) {
48 + if ae, ok := errors.AsType[*apiError](err); ok {
49 utils.WriteAPIError(w, ae.status, ae.code, ae.msg)
50 return
51 }
@@ -138,52 +135,6 @@ func (s *Server) handleHealthz(w http.ResponseWriter, _ *http.Request) {
135 utils.WriteAPIData(w, http.StatusOK, map[string]any{"status": "ok"})
136 }
137
141 -func (s *Server) extractAllowedClientIP(w http.ResponseWriter, r *http.Request) (string, bool) {
142 - clientIP := s.registry.policy.ExtractClientIP(r)
143 - if !s.registry.policy.IPFilter().IsIPBanned(clientIP) {
144 - return clientIP, true
145 - }
146 - utils.WriteAPIError(w, http.StatusForbidden, types.APIErrorCodeIPBanned, "request denied because source IP is banned")
147 - return "", false
148 -}
149 -
150 -func (s *Server) signedRelayDescriptor(now time.Time) (types.RelayDescriptor, error) {
151 - if now.IsZero() {
152 - now = time.Now().UTC()
153 - } else {
154 - now = now.UTC()
155 - }
156 -
157 - var wireGuardPublicKey string
158 - var wireGuardPort int
159 - if s.overlay != nil {
160 - cfg := s.overlay.Config()
161 - wireGuardPublicKey = cfg.PublicKey
162 - wireGuardPort = cfg.ListenPort
163 - }
164 -
165 - self := types.RelayDescriptor{
166 - Address: s.identity.Address,
167 - Version: types.DiscoveryVersion,
168 - IssuedAt: now,
169 - ExpiresAt: now.Add(discovery.DiscoveryDescriptorTTL),
170 - APIHTTPSAddr: s.cfg.PortalURL,
171 - WireGuardPublicKey: wireGuardPublicKey,
172 - WireGuardPort: wireGuardPort,
173 - SupportsOverlay: s.overlay != nil,
174 - SupportsUDP: s.cfg.UDPEnabled && s.quicBackhaul != nil,
175 - SupportsTCP: s.cfg.TCPEnabled,
176 - ActiveConnections: s.proxy.activeConnectionCount(),
177 - TCPBPS: s.proxy.currentTCPBPS(now),
178 - }
179 -
180 - signedSelf, err := auth.SignRelayDescriptor(self, s.identity.PrivateKey)
181 - if err != nil {
182 - return types.RelayDescriptor{}, err
183 - }
184 - return signedSelf, nil
185 -}
186 -
138 func (s *Server) handleRelayDiscovery(w http.ResponseWriter, r *http.Request) {
139 if !utils.RequireMethod(w, r, http.MethodGet) {
140 return
@@ -194,7 +145,7 @@ func (s *Server) handleRelayDiscovery(w http.ResponseWriter, r *http.Request) {
145 }
146
147 now := time.Now().UTC()
197 - self, err := s.signedRelayDescriptor(now)
148 + self, err := s.newSelfDescriptor(now)
149 if err != nil {
150 utils.WriteAPIError(w, http.StatusInternalServerError, types.APIErrorCodeInternal, err.Error())
151 return
@@ -321,13 +272,19 @@ func (s *Server) handleRegister(w http.ResponseWriter, r *http.Request) {
272 return
273 }
274
324 - resp, err := s.registerLease(challenge.Request, clientIP, req.ReportedIP)
275 + record, resp, err := s.registry.Register(challenge.Request, clientIP, req.ReportedIP)
276 if err != nil {
326 - if errors.Is(err, transport.ErrPortExhausted) {
327 - utils.WriteAPIError(w, http.StatusServiceUnavailable, types.APIErrorCodeUDPPortExhausted, err.Error())
328 - } else {
329 - writeAPIErrorResponse(w, err)
277 + writeAPIErrorResponse(w, err)
278 + return
279 + }
280 + if err := s.syncENSGaslessHostname(context.Background(), record); err != nil {
281 + removed, _ := s.registry.Unregister(types.UnregisterRequest{AccessToken: resp.AccessToken})
282 + if removed == nil {
283 + record.Close()
284 + removed = record
285 }
286 + s.deleteENSGaslessHostname(context.Background(), removed, "delete lease ens gasless hostname after sync failure")
287 + writeAPIErrorResponse(w, err)
288 return
289 }
290
@@ -362,8 +319,7 @@ func (s *Server) handleRegisterChallenge(w http.ResponseWriter, r *http.Request)
319 Path: types.PathSDKRegister,
320 }).String()
321
365 - req.HopToken = strings.TrimSpace(req.HopToken)
366 - if req.HopToken != "" && s.hopMux == nil {
322 + if strings.TrimSpace(req.HopToken) != "" && s.hopMux == nil {
323 utils.WriteAPIError(w, http.StatusServiceUnavailable, types.APIErrorCodeFeatureUnavailable, errFeatureUnavailable.Error())
324 return
325 }
@@ -400,31 +356,13 @@ func (s *Server) handleRenew(w http.ResponseWriter, r *http.Request) {
356 return
357 }
358
403 - claims, err := auth.VerifyLeaseAccessToken(req.AccessToken, s.identity.PublicKey, s.cfg.PortalURL, time.Now().UTC())
404 - if err != nil {
405 - utils.WriteAPIError(w, http.StatusForbidden, types.APIErrorCodeUnauthorized, errUnauthorized.Error())
406 - return
407 - }
408 -
409 - ttl := defaultLeaseTTL
410 - if req.TTL > 0 {
411 - ttl = time.Duration(req.TTL) * time.Second
412 - }
413 - record, err := s.registry.Renew(claims.Identity.Key(), ttl, clientIP, utils.SanitizeReportedIP(req.ReportedIP))
359 + resp, err := s.registry.Renew(req, clientIP)
360 if err != nil {
361 writeAPIErrorResponse(w, err)
362 return
363 }
418 - nextAccessToken, _, err := auth.IssueLeaseAccessToken(s.identity.PrivateKey, s.identity.Address, s.cfg.PortalURL, record.Identity, ttl)
419 - if err != nil {
420 - utils.WriteAPIError(w, http.StatusInternalServerError, types.APIErrorCodeInternal, err.Error())
421 - return
422 - }
364
424 - utils.WriteAPIData(w, http.StatusOK, types.RenewResponse{
425 - ExpiresAt: record.ExpiresAt,
426 - AccessToken: nextAccessToken,
427 - })
365 + utils.WriteAPIData(w, http.StatusOK, resp)
366 }
367
368 func (s *Server) handleUnregister(w http.ResponseWriter, r *http.Request) {
@@ -436,41 +374,16 @@ func (s *Server) handleUnregister(w http.ResponseWriter, r *http.Request) {
374 if !ok {
375 return
376 }
439 - claims, err := auth.VerifyLeaseAccessToken(req.AccessToken, s.identity.PublicKey, s.cfg.PortalURL, time.Now().UTC())
440 - if err != nil {
441 - utils.WriteAPIError(w, http.StatusForbidden, types.APIErrorCodeUnauthorized, errUnauthorized.Error())
442 - return
443 - }
444 -
445 - record, err := s.registry.Unregister(claims.Identity.Key())
377 + record, err := s.registry.Unregister(req)
378 if err != nil {
379 writeAPIErrorResponse(w, err)
380 return
381 }
450 - s.cleanupRemovedRecord(context.Background(), record, "delete lease remote state")
382 + s.deleteENSGaslessHostname(context.Background(), record, "delete lease ens gasless hostname")
383
384 utils.WriteAPIData(w, http.StatusOK, map[string]any{})
385 }
386
455 -func (s *Server) cleanupRemovedRecord(ctx context.Context, record *leaseRecord, logMessage string) {
456 - if record == nil {
457 - return
458 - }
459 - if record.isPublicEntry() && s.acmeManager != nil {
460 - deleteCtx, cancel := context.WithTimeout(ctx, defaultClaimTimeout)
461 - err := s.acmeManager.DeleteENSGaslessHostname(deleteCtx, record.Hostname)
462 - cancel()
463 - if err != nil {
464 - log.Warn().
465 - Err(err).
466 - Str("hostname", record.Hostname).
467 - Str("address", record.Address).
468 - Msg(logMessage)
469 - }
470 - }
471 - record.Close()
472 -}
473 -
387 func (s *Server) handleHop(w http.ResponseWriter, r *http.Request) {
388 switch r.Method {
389 case http.MethodPost, http.MethodDelete:
@@ -505,7 +418,7 @@ func (s *Server) handleHop(w http.ResponseWriter, r *http.Request) {
418 }
419 if r.Method == http.MethodDelete {
420 record := s.registry.DeleteHopRoute(&route)
508 - s.cleanupRemovedRecord(context.Background(), record, "delete hop route remote state")
421 + s.deleteENSGaslessHostname(context.Background(), record, "delete hop route ens gasless hostname")
422 utils.WriteAPIData(w, http.StatusOK, map[string]any{})
423 return
424 }
@@ -538,19 +451,14 @@ func (s *Server) handleHop(w http.ResponseWriter, r *http.Request) {
451 writeAPIErrorResponse(w, err)
452 return
453 }
541 - if record.isPublicEntry() && s.acmeManager != nil {
542 - syncCtx, cancel := context.WithTimeout(context.Background(), defaultClaimTimeout)
543 - if err := s.acmeManager.SyncENSGaslessHostname(syncCtx, record.Hostname, record.Address); err != nil {
544 - cancel()
545 - removed := s.registry.DeleteHopRoute(&route)
546 - if removed == nil {
547 - removed = record
548 - }
549 - s.cleanupRemovedRecord(context.Background(), removed, "delete hop route remote state after sync failure")
550 - writeAPIErrorResponse(w, err)
551 - return
454 + if err := s.syncENSGaslessHostname(context.Background(), record); err != nil {
455 + removed := s.registry.DeleteHopRoute(&route)
456 + if removed == nil {
457 + removed = record
458 }
553 - cancel()
459 + s.deleteENSGaslessHostname(context.Background(), removed, "delete hop route ens gasless hostname after sync failure")
460 + writeAPIErrorResponse(w, err)
461 + return
462 }
463 utils.WriteAPIData(w, http.StatusOK, map[string]any{})
464 }
@@ -570,7 +478,7 @@ func (s *Server) handleConnect(w http.ResponseWriter, r *http.Request) {
478 return
479 }
480
573 - lease, err := s.admitLeaseByToken(token, false)
481 + lease, err := s.registry.admitLeaseByToken(token, false)
482 if err != nil {
483 writeAPIErrorResponse(w, err)
484 return
@@ -620,200 +528,11 @@ func (s *Server) handleConnect(w http.ResponseWriter, r *http.Request) {
528 Msg("sdk reverse connected")
529 }
530
623 -func (s *Server) handleQUICBackhaulConn(conn *quic.Conn) {
624 - control, err := transport.AcceptQUICBackhaulControl(context.Background(), conn)
625 - if err != nil {
626 - _ = conn.CloseWithError(1, "control read failed")
627 - return
628 - }
629 -
630 - lease, err := s.admitLeaseByToken(control.AccessToken, true)
631 - if err != nil {
632 - code, reason := types.APIErrorCodeInvalidRequest, "invalid control message"
633 - switch {
634 - case errors.Is(err, errLeaseNotFound):
635 - code, reason = types.APIErrorCodeLeaseNotFound, "lease not found"
636 - case errors.Is(err, errLeaseRejected):
637 - code, reason = types.APIErrorCodeLeaseRejected, "lease rejected"
638 - case errors.Is(err, errUnauthorized):
639 - code, reason = types.APIErrorCodeUnauthorized, "unauthorized"
640 - case errors.Is(err, errTransportMismatch):
641 - code, reason = types.APIErrorCodeTransportMismatch, "transport mismatch"
642 - }
643 - _ = control.Reject(code, reason)
644 - return
645 - }
646 -
647 - if err := lease.datagram.BindBackhaul(conn); err != nil {
648 - _ = control.Reject("broker_closed", "broker closed")
649 - return
650 - }
651 -
652 - _ = control.Accept()
653 - s.registry.Touch(lease.Key(), conn.RemoteAddr().String(), time.Now())
654 - log.Info().
655 - Str("component", "quic-backhaul-listener").
656 - Str("address", lease.Address).
657 - Str("lease_name", lease.Name).
658 - Str("remote_addr", conn.RemoteAddr().String()).
659 - Msg("quic backhaul connected")
660 -}
661 -
662 -func (s *Server) admitLeaseByToken(token string, requireDatagram bool) (*leaseRecord, error) {
663 - claims, err := auth.VerifyLeaseAccessToken(token, s.identity.PublicKey, s.cfg.PortalURL, time.Now().UTC())
664 - if err != nil {
665 - return nil, errUnauthorized
666 - }
667 - lease, ok := s.registry.RecordByKey(claims.Identity.Key(), time.Now())
668 - if !ok {
669 - return nil, errLeaseNotFound
670 - }
671 - if !s.registry.policy.IsIdentityRoutable(lease.Key()) {
672 - return nil, errLeaseRejected
673 - }
674 - if lease.stream == nil || (requireDatagram && lease.datagram == nil) {
675 - return nil, errTransportMismatch
676 - }
677 - return lease, nil
678 -}
679 -
680 -func (s *Server) registerLease(req types.RegisterChallengeRequest, clientIP, reportedIP string) (types.RegisterResponse, error) {
681 - identity, err := utils.NormalizeIdentity(req.Identity)
682 - if err != nil {
683 - return types.RegisterResponse{}, err
684 - }
685 - if s.registry.policy.IPFilter().IsIPBanned(clientIP) {
686 - return types.RegisterResponse{}, errIPBanned
687 - }
688 - hostname, err := utils.LeaseHostname(identity.Name, s.identity.Name)
689 - if err != nil {
690 - return types.RegisterResponse{}, err
691 - }
692 -
693 - ttl := defaultLeaseTTL
694 - if req.TTL > 0 {
695 - ttl = time.Duration(req.TTL) * time.Second
696 - }
697 -
698 - if req.UDPEnabled {
699 - if !s.cfg.UDPEnabled || s.group != nil && s.quicBackhaul == nil {
700 - return types.RegisterResponse{}, errFeatureUnavailable
701 - }
702 - if !s.registry.policy.IsUDPEnabled() {
703 - return types.RegisterResponse{}, errUDPDisabled
704 - }
705 - if max := s.registry.policy.UDPMaxLeases(); max > 0 && s.registry.countDatagramLeases() >= max {
706 - return types.RegisterResponse{}, errUDPCapacityExceeded
707 - }
708 - }
709 - if req.TCPEnabled {
710 - if !s.cfg.TCPEnabled {
711 - return types.RegisterResponse{}, errFeatureUnavailable
712 - }
713 - if !s.registry.policy.IsTCPPortEnabled() {
714 - return types.RegisterResponse{}, errTCPPortDisabled
715 - }
716 - if max := s.registry.policy.TCPPortMaxLeases(); max > 0 && s.registry.countTCPPortLeases() >= max {
717 - return types.RegisterResponse{}, errTCPPortCapacityExceeded
718 - }
719 - }
720 - accessToken, claims, err := auth.IssueLeaseAccessToken(s.identity.PrivateKey, s.identity.Address, s.cfg.PortalURL, identity, ttl)
721 - if err != nil {
722 - return types.RegisterResponse{}, err
723 - }
724 - issuedAt := claims.IssuedAt.Time().UTC()
725 - expiresAt := claims.Expiry.Time().UTC()
726 - req.HopToken = strings.TrimSpace(req.HopToken)
727 - if req.HopToken != "" && s.hopMux == nil {
728 - return types.RegisterResponse{}, errFeatureUnavailable
729 - }
730 - identityKey := identity.Key()
731 - stream := transport.NewRelayStream(identityKey, defaultIdleKeepalive, defaultReadyQueueLimit)
732 - record := &leaseRecord{
733 - Identity: identity,
734 - Hostname: hostname,
735 - Metadata: req.Metadata,
736 - ExpiresAt: expiresAt,
737 - FirstSeenAt: issuedAt,
738 - LastSeenAt: issuedAt,
739 - ClientIP: clientIP,
740 - ReportedIP: utils.SanitizeReportedIP(reportedIP),
741 - hopToken: req.HopToken,
742 - stream: stream,
743 - }
744 - if req.UDPEnabled {
745 - if s.udpPorts == nil {
746 - return types.RegisterResponse{}, errors.New("udp port allocation not available")
747 - }
748 - port, err := s.udpPorts.Allocate(identity.Name)
749 - if err != nil {
750 - return types.RegisterResponse{}, err
751 - }
752 - record.datagram = transport.NewRelayDatagram(identityKey, port)
753 - record.udpPorts = s.udpPorts
754 - }
755 - if req.TCPEnabled {
756 - if s.tcpPorts == nil {
757 - return types.RegisterResponse{}, errors.New("tcp port allocation not available")
758 - }
759 - port, err := s.tcpPorts.Allocate(identity.Name)
760 - if err != nil {
761 - if errors.Is(err, transport.ErrPortExhausted) {
762 - return types.RegisterResponse{}, errTCPPortExhausted
763 - }
764 - return types.RegisterResponse{}, err
765 - }
766 - record.tcpPort = transport.NewRelayTCPPort(identityKey, port, stream, func(left, right net.Conn) {
767 - s.proxy.bridge(left, right, identityKey, s.registry.policy.BPSManager())
768 - })
769 - record.tcpPorts = s.tcpPorts
770 - }
771 -
772 - if err := record.Start(); err != nil {
773 - record.Close()
774 - return types.RegisterResponse{}, err
775 - }
776 -
777 - if err := s.registry.Register(record); err != nil {
778 - record.Close()
779 - return types.RegisterResponse{}, err
780 - }
781 - if record.isPublicEntry() {
782 - syncCtx, cancel := context.WithTimeout(context.Background(), defaultClaimTimeout)
783 - defer cancel()
784 - if err := s.acmeManager.SyncENSGaslessHostname(syncCtx, record.Hostname, record.Address); err != nil {
785 - removed, _ := s.registry.Unregister(record.Key())
786 - if removed == nil {
787 - removed = record
788 - }
789 - s.cleanupRemovedRecord(context.Background(), removed, "delete lease remote state after sync failure")
790 - return types.RegisterResponse{}, err
791 - }
792 - }
793 -
794 - resp := types.RegisterResponse{
795 - Identity: record.Identity,
796 - Hostname: hostname,
797 - ExpiresAt: expiresAt,
798 - AccessToken: accessToken,
799 - UDPEnabled: record.datagram != nil,
800 - TCPEnabled: record.tcpPort != nil,
801 - }
802 - if record.datagram != nil {
803 - resp.SNIPort = s.cfg.SNIPort
804 - resp.UDPAddr = fmt.Sprintf("%s:%d", s.identity.Name, record.datagram.UDPPort())
805 - }
806 - if record.tcpPort != nil {
807 - resp.TCPAddr = fmt.Sprintf("%s:%d", s.identity.Name, record.tcpPort.TCPPort())
808 - }
809 -
810 - return resp, nil
811 -}
812 -
813 -func (s *Server) runAPIServer() error {
814 - err := s.apiServer.Serve(s.apiListener)
815 - if err == nil || errors.Is(err, http.ErrServerClosed) || errors.Is(err, net.ErrClosed) {
816 - return nil
531 +func (s *Server) extractAllowedClientIP(w http.ResponseWriter, r *http.Request) (string, bool) {
532 + clientIP := s.registry.policy.ExtractClientIP(r)
533 + if !s.registry.policy.IPFilter().IsIPBanned(clientIP) {
534 + return clientIP, true
535 }
818 - return err
536 + utils.WriteAPIError(w, http.StatusForbidden, types.APIErrorCodeIPBanned, "request denied because source IP is banned")
537 + return "", false
538 }
portal/lease.go
+259 -101
@@ -4,6 +4,8 @@ import (
4 "context"
5 "errors"
6 "fmt"
7 + "net"
8 + "net/http"
9 "strings"
10 "sync"
11 "time"
@@ -15,30 +17,52 @@ import (
17 "github.com/gosuda/portal-tunnel/v2/utils"
18 )
19
18 -const defaultRegisterChallengeTTL = 2 * time.Minute
20 +const (
21 + defaultLeaseTTL = 30 * time.Second
22 + defaultRegisterChallengeTTL = 2 * time.Minute
23 + defaultPortReservationGrace = 5 * time.Minute
24 + defaultIdleKeepalive = 15 * time.Second
25 + defaultReadyQueueLimit = 8
26 +)
27
28 type leaseRegistry struct {
21 - records []*leaseRecord
22 - policy *policy.Runtime
23 - mu sync.RWMutex
29 + records []*leaseRecord
30 + rootHostname string
31 + sniPort int
32 + tokenPrivateKey string
33 + tokenPublicKey string
34 + tokenKeyID string
35 + tokenIssuer string
36 + policy *policy.Runtime
37 + udpPorts *transport.PortAllocator
38 + tcpPorts *transport.PortAllocator
39 + proxy *proxy
40 + mu sync.RWMutex
41 }
42
26 -func newLeaseRegistry(udpEnabled, tcpPortEnabled bool, trustProxyHeaders bool, rawTrustedProxyCIDRs string) (*leaseRegistry, error) {
43 +func newLeaseRegistry(udpEnabled, tcpPortEnabled bool, minPort, maxPort int, rootHostname string, sniPort int, tokenPrivateKey, tokenPublicKey, tokenKeyID, tokenIssuer string, trustProxyHeaders bool, rawTrustedProxyCIDRs string) (*leaseRegistry, error) {
44 runtime, err := policy.NewRuntime(udpEnabled, tcpPortEnabled, trustProxyHeaders, rawTrustedProxyCIDRs)
45 if err != nil {
46 return nil, err
47 }
48
49 return &leaseRegistry{
33 - records: make([]*leaseRecord, 0),
34 - policy: runtime,
50 + records: make([]*leaseRecord, 0),
51 + rootHostname: utils.NormalizeHostname(rootHostname),
52 + sniPort: sniPort,
53 + tokenPrivateKey: tokenPrivateKey,
54 + tokenPublicKey: tokenPublicKey,
55 + tokenKeyID: tokenKeyID,
56 + tokenIssuer: tokenIssuer,
57 + policy: runtime,
58 + udpPorts: transport.NewPortAllocator(minPort, maxPort, defaultPortReservationGrace),
59 + tcpPorts: transport.NewPortAllocator(minPort, maxPort, defaultPortReservationGrace),
60 + proxy: &proxy{},
61 }, nil
62 }
63
64 func (r *leaseRegistry) CloseAll() []*leaseRecord {
65 r.mu.Lock()
40 - defer r.mu.Unlock()
41 -
66 out := r.records
67 for _, record := range out {
68 if record != nil && record.stream != nil {
@@ -46,6 +70,11 @@ func (r *leaseRegistry) CloseAll() []*leaseRecord {
70 }
71 }
72 r.records = nil
73 + r.mu.Unlock()
74 +
75 + for _, record := range out {
76 + record.Close()
77 + }
78 return out
79 }
80
@@ -102,82 +131,223 @@ func (r *leaseRegistry) recordByHopToken(token string, now time.Time) *leaseReco
131 return nil
132 }
133
105 -func (r *leaseRegistry) Register(record *leaseRecord) error {
106 - if record == nil {
107 - return errors.New("lease record is required")
134 +func (r *leaseRegistry) Register(req types.RegisterChallengeRequest, clientIP, reportedIP string) (*leaseRecord, types.RegisterResponse, error) {
135 + if r == nil {
136 + return nil, types.RegisterResponse{}, errFeatureUnavailable
137 + }
138 + identity, err := utils.NormalizeIdentity(req.Identity)
139 + if err != nil {
140 + return nil, types.RegisterResponse{}, err
141 + }
142 + if r.policy.IPFilter().IsIPBanned(clientIP) {
143 + return nil, types.RegisterResponse{}, errIPBanned
144 }
145
110 - key := record.Key()
111 - if key == "" {
112 - return errors.New("lease identity is required")
146 + ttl := defaultLeaseTTL
147 + if req.TTL > 0 {
148 + ttl = time.Duration(req.TTL) * time.Second
149 }
114 - hostname := utils.NormalizeHostname(record.Hostname)
115 - if hostname == "" {
116 - return errors.New("lease hostname is required")
150 +
151 + identityKey := identity.Key()
152 + hostname, err := utils.LeaseHostname(identity.Name, r.rootHostname)
153 + if err != nil {
154 + return nil, types.RegisterResponse{}, err
155 + }
156 + hopToken := strings.TrimSpace(req.HopToken)
157 + if hopToken != "" && (req.UDPEnabled || req.TCPEnabled) {
158 + return nil, types.RegisterResponse{}, errTransportMismatch
159 + }
160 + if req.UDPEnabled {
161 + if !r.policy.IsUDPEnabled() {
162 + return nil, types.RegisterResponse{}, errUDPDisabled
163 + }
164 + if max := r.policy.UDPMaxLeases(); max > 0 && r.countDatagramLeases() >= max {
165 + return nil, types.RegisterResponse{}, errUDPCapacityExceeded
166 + }
167 + }
168 + if req.TCPEnabled {
169 + if !r.policy.IsTCPPortEnabled() {
170 + return nil, types.RegisterResponse{}, errTCPPortDisabled
171 + }
172 + if max := r.policy.TCPPortMaxLeases(); max > 0 && r.countTCPPortLeases() >= max {
173 + return nil, types.RegisterResponse{}, errTCPPortCapacityExceeded
174 + }
175 + if r.proxy == nil {
176 + return nil, types.RegisterResponse{}, errors.New("tcp proxy is not available")
177 + }
178 }
118 - record.hopToken = strings.TrimSpace(record.hopToken)
179
120 - r.mu.Lock()
180 + accessToken, claims, err := auth.IssueLeaseAccessToken(r.tokenPrivateKey, r.tokenKeyID, r.tokenIssuer, identity, ttl)
181 + if err != nil {
182 + return nil, types.RegisterResponse{}, err
183 + }
184 + issuedAt := claims.IssuedAt.Time().UTC()
185 + expiresAt := claims.Expiry.Time().UTC()
186
122 - now := time.Now()
123 - if record.isPublicEntry() {
124 - for _, existing := range r.records {
125 - if existing == nil || !existing.isPublicEntry() || existing.isExpired(now) {
126 - continue
127 - }
128 - if existing.Hostname == hostname && existing.Key() != key {
129 - r.mu.Unlock()
130 - return errHostnameConflict
187 + stream := transport.NewRelayStream(identityKey, defaultIdleKeepalive, defaultReadyQueueLimit)
188 + record := &leaseRecord{
189 + Identity: identity,
190 + Hostname: hostname,
191 + Metadata: req.Metadata,
192 + ExpiresAt: expiresAt,
193 + FirstSeenAt: issuedAt,
194 + LastSeenAt: issuedAt,
195 + ClientIP: clientIP,
196 + ReportedIP: utils.SanitizeReportedIP(reportedIP),
197 + hopToken: hopToken,
198 + stream: stream,
199 + }
200 +
201 + if req.UDPEnabled {
202 + if r.udpPorts == nil {
203 + return nil, types.RegisterResponse{}, errors.New("udp port allocation not available")
204 + }
205 + port, err := r.udpPorts.Allocate(identity.Name)
206 + if err != nil {
207 + if errors.Is(err, transport.ErrPortExhausted) {
208 + return nil, types.RegisterResponse{}, errUDPPortExhausted
209 }
210 + return nil, types.RegisterResponse{}, err
211 }
212 + record.datagram = transport.NewRelayDatagram(identityKey, port)
213 + record.udpPorts = r.udpPorts
214 }
134 - if record.isHopExit() {
135 - if existing := r.recordByHopToken(record.hopToken, now); existing != nil && existing.Key() != key {
136 - r.mu.Unlock()
137 - return errors.New("hop token conflict")
215 +
216 + if req.TCPEnabled {
217 + if r.tcpPorts == nil {
218 + record.Close()
219 + return nil, types.RegisterResponse{}, errors.New("tcp port allocation not available")
220 + }
221 + port, err := r.tcpPorts.Allocate(identity.Name)
222 + if err != nil {
223 + record.Close()
224 + if errors.Is(err, transport.ErrPortExhausted) {
225 + return nil, types.RegisterResponse{}, errTCPPortExhausted
226 + }
227 + return nil, types.RegisterResponse{}, err
228 }
229 + record.tcpPort = transport.NewRelayTCPPort(identityKey, port, stream, func(left, right net.Conn) {
230 + r.proxy.bridge(left, right, identityKey, r.policy.BPSManager())
231 + })
232 + record.tcpPorts = r.tcpPorts
233 + }
234 +
235 + if err := record.Start(); err != nil {
236 + record.Close()
237 + return nil, types.RegisterResponse{}, err
238 }
239
240 var replaced *leaseRecord
241 replacedIndex := -1
242 + r.mu.Lock()
243 + now := time.Now()
244 + for _, existing := range r.records {
245 + if existing == nil || !existing.isPublicEntry() || existing.isExpired(now) {
246 + continue
247 + }
248 + if existing.Hostname == hostname && existing.Key() != identityKey {
249 + r.mu.Unlock()
250 + record.Close()
251 + return nil, types.RegisterResponse{}, errHostnameConflict
252 + }
253 + }
254 + if hopToken != "" {
255 + if existing := r.recordByHopToken(hopToken, now); existing != nil && existing.Key() != identityKey {
256 + r.mu.Unlock()
257 + record.Close()
258 + return nil, types.RegisterResponse{}, errors.New("hop token conflict")
259 + }
260 + }
261 for i := 0; i < len(r.records); i++ {
262 existing := r.records[i]
263 if existing != nil && existing.stream == nil && existing.isPublicEntry() &&
146 - existing.Hostname == hostname && existing.Key() == key {
264 + existing.Hostname == hostname && existing.Key() == identityKey {
265 r.deleteRecord(i)
266 i--
267 }
268 }
269 for i, existing := range r.records {
152 - if existing != nil && existing.stream != nil && existing.Key() == key {
270 + if existing != nil && existing.stream != nil && existing.Key() == identityKey {
271 replaced = existing
272 replacedIndex = i
155 - r.policy.ForgetIdentity(existing.Key())
273 + r.policy.ForgetIdentity(identityKey)
274 break
275 }
276 }
159 - record.Hostname = hostname
277 if replacedIndex >= 0 {
278 r.records[replacedIndex] = record
279 } else {
280 r.records = append(r.records, record)
281 }
165 - r.policy.IPFilter().RegisterIdentityIP(key, record.ClientIP)
282 + r.policy.IPFilter().RegisterIdentityIP(identityKey, record.ClientIP)
283 r.mu.Unlock()
284
168 - if replaced != nil && replaced != record {
285 + if replaced != nil {
286 replaced.Close()
287 }
171 - return nil
288 +
289 + resp := types.RegisterResponse{
290 + Identity: record.Identity,
291 + Hostname: record.Hostname,
292 + ExpiresAt: record.ExpiresAt,
293 + AccessToken: accessToken,
294 + UDPEnabled: record.datagram != nil,
295 + TCPEnabled: record.tcpPort != nil,
296 + }
297 + if record.datagram != nil {
298 + resp.SNIPort = r.sniPort
299 + resp.UDPAddr = fmt.Sprintf("%s:%d", r.rootHostname, record.datagram.UDPPort())
300 + }
301 + if record.tcpPort != nil {
302 + resp.TCPAddr = fmt.Sprintf("%s:%d", r.rootHostname, record.tcpPort.TCPPort())
303 + }
304 + return record, resp, nil
305 }
306
174 -func (r *leaseRegistry) Renew(key string, ttl time.Duration, clientIP, reportedIP string) (*leaseRecord, error) {
175 - r.mu.Lock()
176 - defer r.mu.Unlock()
307 +func (r *leaseRegistry) admitLeaseByToken(token string, requireDatagram bool) (*leaseRecord, error) {
308 + if r == nil {
309 + return nil, errFeatureUnavailable
310 + }
311 + now := time.Now().UTC()
312 + claims, err := auth.VerifyLeaseAccessToken(token, r.tokenPublicKey, r.tokenIssuer, now)
313 + if err != nil {
314 + return nil, errUnauthorized
315 + }
316 + r.mu.RLock()
317 + lease := r.recordByKey(claims.Identity.Key(), now)
318 + r.mu.RUnlock()
319 + if lease == nil {
320 + return nil, errLeaseNotFound
321 + }
322 + if !r.policy.IsIdentityRoutable(lease.Key()) {
323 + return nil, errLeaseRejected
324 + }
325 + if lease.stream == nil || (requireDatagram && lease.datagram == nil) {
326 + return nil, errTransportMismatch
327 + }
328 + return lease, nil
329 +}
330
178 - record := r.recordByKey(key, time.Time{})
331 +func (r *leaseRegistry) Renew(req types.RenewRequest, clientIP string) (types.RenewResponse, error) {
332 + if r == nil {
333 + return types.RenewResponse{}, errFeatureUnavailable
334 + }
335 + claims, err := auth.VerifyLeaseAccessToken(req.AccessToken, r.tokenPublicKey, r.tokenIssuer, time.Now().UTC())
336 + if err != nil {
337 + return types.RenewResponse{}, errUnauthorized
338 + }
339 + ttl := defaultLeaseTTL
340 + if req.TTL > 0 {
341 + ttl = time.Duration(req.TTL) * time.Second
342 + }
343 +
344 + leaseKey := claims.Identity.Key()
345 + reportedIP := utils.SanitizeReportedIP(req.ReportedIP)
346 + r.mu.Lock()
347 + record := r.recordByKey(leaseKey, time.Time{})
348 if record == nil {
180 - return nil, errLeaseNotFound
349 + r.mu.Unlock()
350 + return types.RenewResponse{}, errLeaseNotFound
351 }
352
353 now := time.Now()
@@ -190,58 +360,46 @@ func (r *leaseRegistry) Renew(key string, ttl time.Duration, clientIP, reportedI
360 if strings.TrimSpace(reportedIP) != "" {
361 record.ReportedIP = reportedIP
362 }
193 - r.policy.IPFilter().RegisterIdentityIP(record.Key(), clientIP)
194 - return record, nil
363 + r.policy.IPFilter().RegisterIdentityIP(leaseKey, clientIP)
364 + identity := record.Identity
365 + r.mu.Unlock()
366 +
367 + nextAccessToken, _, err := auth.IssueLeaseAccessToken(r.tokenPrivateKey, r.tokenKeyID, r.tokenIssuer, identity, ttl)
368 + if err != nil {
369 + return types.RenewResponse{}, &apiError{types.APIErrorCodeInternal, err.Error(), http.StatusInternalServerError}
370 + }
371 +
372 + return types.RenewResponse{
373 + ExpiresAt: expiresAt,
374 + AccessToken: nextAccessToken,
375 + }, nil
376 }
377
197 -func (r *leaseRegistry) Unregister(key string) (*leaseRecord, error) {
378 +func (r *leaseRegistry) Unregister(req types.UnregisterRequest) (*leaseRecord, error) {
379 + if r == nil {
380 + return nil, errFeatureUnavailable
381 + }
382 + claims, err := auth.VerifyLeaseAccessToken(req.AccessToken, r.tokenPublicKey, r.tokenIssuer, time.Now().UTC())
383 + if err != nil {
384 + return nil, errUnauthorized
385 + }
386 r.mu.Lock()
199 - defer r.mu.Unlock()
387
201 - key = strings.TrimSpace(key)
388 + key := strings.TrimSpace(claims.Identity.Key())
389 for i, record := range r.records {
390 if record == nil || record.stream == nil || record.Key() != key {
391 continue
392 }
393 r.deleteRecord(i)
394 r.policy.ForgetIdentity(key)
395 + r.mu.Unlock()
396 + record.Close()
397 return record, nil
398 }
399 + r.mu.Unlock()
400 return nil, errLeaseNotFound
401 }
402
213 -func (r *leaseRegistry) RecordByKey(key string, now time.Time) (*leaseRecord, bool) {
214 - key = strings.TrimSpace(key)
215 - if key == "" {
216 - return nil, false
217 - }
218 -
219 - r.mu.RLock()
220 - defer r.mu.RUnlock()
221 -
222 - record := r.recordByKey(key, now)
223 - if record == nil {
224 - return nil, false
225 - }
226 - return record, true
227 -}
228 -
229 -func (r *leaseRegistry) RecordByHopToken(token string, now time.Time) (*leaseRecord, bool) {
230 - token = strings.TrimSpace(token)
231 - if token == "" {
232 - return nil, false
233 - }
234 -
235 - r.mu.RLock()
236 - defer r.mu.RUnlock()
237 -
238 - record := r.recordByHopToken(token, now)
239 - if record != nil {
240 - return record, true
241 - }
242 - return nil, false
243 -}
244 -
403 func (r *leaseRegistry) RegisterHopRoute(route *types.HopRoute, now time.Time) (*leaseRecord, error) {
404 if route == nil {
405 return nil, errors.New("hop route is required")
@@ -370,6 +528,7 @@ func (r *leaseRegistry) DeleteHopRoute(route *types.HopRoute) *leaseRecord {
528 }
529 }
530 r.mu.Unlock()
531 + deleted.Close()
532 return deleted
533 }
534
@@ -460,7 +619,6 @@ func (r *leaseRegistry) Touch(key, clientIP string, now time.Time) {
619
620 func (r *leaseRegistry) cleanupExpired(now time.Time) []*leaseRecord {
621 r.mu.Lock()
463 - defer r.mu.Unlock()
622
623 var expired []*leaseRecord
624 for i := 0; i < len(r.records); {
@@ -475,6 +633,11 @@ func (r *leaseRegistry) cleanupExpired(now time.Time) []*leaseRecord {
633 }
634 i++
635 }
636 + r.mu.Unlock()
637 +
638 + for _, record := range expired {
639 + record.Close()
640 + }
641 return expired
642 }
643
@@ -611,13 +774,11 @@ type leaseRecord struct {
774 hopNextToken string
775 registerChallenge *auth.RegisterChallenge
776
614 - datagram *transport.RelayDatagram
615 - udpPorts *transport.PortAllocator
616 - tcpPort *transport.RelayTCPPort
617 - tcpPorts *transport.PortAllocator
618 - stream *transport.RelayStream
619 - startErr error
620 - startOnce sync.Once
777 + datagram *transport.RelayDatagram
778 + udpPorts *transport.PortAllocator
779 + tcpPort *transport.RelayTCPPort
780 + tcpPorts *transport.PortAllocator
781 + stream *transport.RelayStream
782 }
783
784 func (r *leaseRecord) isPublicEntry() bool {
@@ -648,18 +809,15 @@ func (r *leaseRecord) isExpired(now time.Time) bool {
809 }
810
811 func (r *leaseRecord) Start() error {
651 - r.startOnce.Do(func() {
652 - if r.datagram != nil {
653 - r.startErr = r.datagram.Start(context.Background())
654 - if r.startErr != nil {
655 - return
656 - }
657 - }
658 - if r.tcpPort != nil {
659 - r.startErr = r.tcpPort.Start(context.Background())
812 + if r.datagram != nil {
813 + if err := r.datagram.Start(context.Background()); err != nil {
814 + return err
815 }
661 - })
662 - return r.startErr
816 + }
817 + if r.tcpPort != nil {
818 + return r.tcpPort.Start(context.Background())
819 + }
820 + return nil
821 }
822
823 func (r *leaseRecord) Close() {
portal/lease_test.go
+53 -63
@@ -10,33 +10,41 @@ import (
10 "github.com/gosuda/portal-tunnel/v2/portal/policy"
11 "github.com/gosuda/portal-tunnel/v2/portal/transport"
12 "github.com/gosuda/portal-tunnel/v2/types"
13 + "github.com/gosuda/portal-tunnel/v2/utils"
14 )
15
16 func newTestRegistry(t *testing.T) *leaseRegistry {
17 t.Helper()
17 - registry, err := newLeaseRegistry(false, false, false, "")
18 + relay, err := utils.LoadOrCreateRelayIdentity(t.TempDir(), "example.com", false)
19 + if err != nil {
20 + t.Fatalf("LoadOrCreateRelayIdentity() error = %v", err)
21 + }
22 + registry, err := newLeaseRegistry(false, false, 10000, 10100, relay.Name, 443, relay.PrivateKey, relay.PublicKey, relay.Address, "https://example.com", false, "")
23 if err != nil {
24 t.Fatalf("newLeaseRegistry() error = %v", err)
25 }
26 return registry
27 }
28
29 +func newTestLeaseIdentity(t *testing.T, name string) types.Identity {
30 + t.Helper()
31 + identity, err := utils.ResolveSecp256k1Identity("")
32 + if err != nil {
33 + t.Fatalf("ResolveSecp256k1Identity() error = %v", err)
34 + }
35 + identity.Name = name
36 + return identity
37 +}
38 +
39 func TestLeaseRegistryLifecycle(t *testing.T) {
40 t.Parallel()
41
42 registry := newTestRegistry(t)
43 runtime := registry.policy
29 - record := &leaseRecord{
30 - Identity: types.Identity{
31 - Name: "demo",
32 - Address: "addr-1",
33 - },
34 - Hostname: "demo.example.com",
35 - ExpiresAt: time.Now().Add(30 * time.Second),
36 - stream: transport.NewRelayStream("addr-1", time.Minute, 1),
37 - }
38 -
39 - if err := registry.Register(record); err != nil {
44 + record, registered, err := registry.Register(types.RegisterChallengeRequest{
45 + Identity: newTestLeaseIdentity(t, "demo"),
46 + }, "203.0.113.10", "")
47 + if err != nil {
48 t.Fatalf("Register() error = %v", err)
49 }
50
@@ -45,18 +53,27 @@ func TestLeaseRegistryLifecycle(t *testing.T) {
53 t.Fatalf("Lookup() = %v, %v, want registered lease", lookedUp, ok)
54 }
55
48 - renewed, err := registry.Renew(record.Key(), time.Minute, "203.0.113.10", "")
56 + renewed, err := registry.Renew(types.RenewRequest{
57 + AccessToken: registered.AccessToken,
58 + TTL: int(time.Minute / time.Second),
59 + }, "203.0.113.11")
60 if err != nil {
61 t.Fatalf("Renew() error = %v", err)
62 }
52 - if renewed.ClientIP != "203.0.113.10" {
53 - t.Fatalf("Renew() client ip = %q, want %q", renewed.ClientIP, "203.0.113.10")
63 + if record.ClientIP != "203.0.113.11" {
64 + t.Fatalf("Renew() client ip = %q, want %q", record.ClientIP, "203.0.113.11")
65 + }
66 + if !renewed.ExpiresAt.Equal(record.ExpiresAt) {
67 + t.Fatalf("Renew() expires at = %v, want %v", renewed.ExpiresAt, record.ExpiresAt)
68 }
55 - if got := runtime.IPFilter().IdentityIP(record.Key()); got != "203.0.113.10" {
69 + if renewed.AccessToken == "" {
70 + t.Fatal("Renew() access token is empty")
71 + }
72 + if got := runtime.IPFilter().IdentityIP(record.Key()); got != "203.0.113.11" {
73 t.Fatalf("Renew() did not register client IP for lease")
74 }
75
59 - removed, err := registry.Unregister(record.Key())
76 + removed, err := registry.Unregister(types.UnregisterRequest{AccessToken: renewed.AccessToken})
77 if err != nil {
78 t.Fatalf("Unregister() error = %v", err)
79 }
@@ -77,17 +94,11 @@ func TestLeaseRegistryWildcardAndConflict(t *testing.T) {
94
95 registry := newTestRegistry(t)
96 wildcardLease := &leaseRecord{
80 - Identity: types.Identity{
81 - Name: "wildcard",
82 - Address: "addr-wildcard",
83 - },
97 + Identity: newTestLeaseIdentity(t, "wildcard"),
98 Hostname: "*.example.com",
99 ExpiresAt: time.Now().Add(30 * time.Second),
86 - stream: transport.NewRelayStream("addr-wildcard", time.Minute, 1),
87 - }
88 - if err := registry.Register(wildcardLease); err != nil {
89 - t.Fatalf("Register(wildcard) error = %v", err)
100 }
101 + registry.records = append(registry.records, wildcardLease)
102
103 if _, ok := registry.Lookup("app.example.com"); !ok {
104 t.Fatal("Lookup(one-level wildcard) = false, want true")
@@ -96,18 +107,16 @@ func TestLeaseRegistryWildcardAndConflict(t *testing.T) {
107 t.Fatal("Lookup(multi-level wildcard) = true, want false")
108 }
109
99 - conflict := &leaseRecord{
100 - Identity: types.Identity{
101 - Name: "conflict",
102 - Address: "addr-conflict",
103 - },
104 - Hostname: "*.example.com",
105 - ExpiresAt: time.Now().Add(30 * time.Second),
106 - stream: transport.NewRelayStream("addr-conflict", time.Minute, 1),
110 + if _, _, err := registry.Register(types.RegisterChallengeRequest{
111 + Identity: newTestLeaseIdentity(t, "conflict"),
112 + }, "203.0.113.10", ""); err != nil {
113 + t.Fatalf("Register(conflict first) error = %v", err)
114 }
108 - err := registry.Register(conflict)
115 + _, _, err := registry.Register(types.RegisterChallengeRequest{
116 + Identity: newTestLeaseIdentity(t, "conflict"),
117 + }, "203.0.113.11", "")
118 if !errors.Is(err, errHostnameConflict) {
110 - t.Fatalf("Register(conflict) error = %v, want hostname conflict", err)
119 + t.Fatalf("Register(conflict second) error = %v, want hostname conflict", err)
120 }
121 }
122
@@ -119,17 +128,10 @@ func TestLeaseRegistryAdminLeasesAndRoutableUsePolicy(t *testing.T) {
128 if err := runtime.Approver().SetMode(policy.ModeManual); err != nil {
129 t.Fatalf("SetMode() error = %v", err)
130 }
122 - record := &leaseRecord{
123 - Identity: types.Identity{
124 - Name: "demo",
125 - Address: "addr-policy",
126 - },
127 - Hostname: "demo.example.com",
128 - ExpiresAt: time.Now().Add(30 * time.Second),
129 - ClientIP: "203.0.113.20",
130 - stream: transport.NewRelayStream("addr-policy", time.Minute, 1),
131 - }
132 - if err := registry.Register(record); err != nil {
131 + record, _, err := registry.Register(types.RegisterChallengeRequest{
132 + Identity: newTestLeaseIdentity(t, "demo"),
133 + }, "203.0.113.20", "")
134 + if err != nil {
135 t.Fatalf("Register() error = %v", err)
136 }
137
@@ -170,16 +172,11 @@ func TestLeaseRegistryPublicLeasesIncludesIngressRouteInManualApproval(t *testin
172 t.Fatalf("SetMode() error = %v", err)
173 }
174 route := &leaseRecord{
173 - Identity: types.Identity{
174 - Name: "demo",
175 - Address: "addr-ingress",
176 - },
175 + Identity: newTestLeaseIdentity(t, "demo"),
176 Hostname: "demo.example.com",
177 ExpiresAt: time.Now().Add(30 * time.Second),
178 }
180 - if err := registry.Register(route); err != nil {
181 - t.Fatalf("Register() error = %v", err)
182 - }
179 + registry.records = append(registry.records, route)
180
181 leases := registry.PublicLeases(time.Now())
182 if len(leases) != 1 {
@@ -195,21 +192,14 @@ func TestLeaseRegistryCleanupExpiredClosesBroker(t *testing.T) {
192
193 registry := newTestRegistry(t)
194 record := &leaseRecord{
198 - Identity: types.Identity{
199 - Name: "expired",
200 - Address: "addr-expired",
201 - },
195 + Identity: newTestLeaseIdentity(t, "expired"),
196 Hostname: "expired.example.com",
197 ExpiresAt: time.Now().Add(-time.Second),
198 stream: transport.NewRelayStream("addr-expired", time.Minute, 1),
199 }
206 - if err := registry.Register(record); err != nil {
207 - t.Fatalf("Register() error = %v", err)
208 - }
200 + registry.records = append(registry.records, record)
201
210 - for _, lease := range registry.cleanupExpired(time.Now()) {
211 - lease.Close()
212 - }
202 + registry.cleanupExpired(time.Now())
203
204 if _, ok := registry.Lookup("expired.example.com"); ok {
205 t.Fatal("Lookup() after cleanupExpired() = true, want false")
portal/server.go
+116 -15
@@ -18,6 +18,7 @@ import (
18 "golang.org/x/sync/errgroup"
19
20 "github.com/gosuda/portal-tunnel/v2/portal/acme"
21 + "github.com/gosuda/portal-tunnel/v2/portal/auth"
22 "github.com/gosuda/portal-tunnel/v2/portal/discovery"
23 "github.com/gosuda/portal-tunnel/v2/portal/keyless"
24 "github.com/gosuda/portal-tunnel/v2/portal/overlay"
@@ -28,10 +29,7 @@ import (
29 )
30
31 const (
31 - defaultLeaseTTL = 30 * time.Second
32 defaultClaimTimeout = 10 * time.Second
33 - defaultIdleKeepalive = 15 * time.Second
34 - defaultReadyQueueLimit = 8
33 defaultClientHelloWait = 2 * time.Second
34 defaultControlBodyLimit = 4 << 20
35 defaultHopOpenRetryWait = 250 * time.Millisecond
@@ -123,8 +121,6 @@ type Server struct {
121 relaySet *discovery.RelaySet
122 announceLimiter *discovery.AnnounceLimiter
123 registry *leaseRegistry
126 - udpPorts *transport.PortAllocator
127 - tcpPorts *transport.PortAllocator
124 }
125
126 func NewServer(cfg ServerConfig) (*Server, error) {
@@ -137,7 +133,7 @@ func NewServer(cfg ServerConfig) (*Server, error) {
133 if err != nil {
134 return nil, fmt.Errorf("load relay identity: %w", err)
135 }
140 - registry, err := newLeaseRegistry(cfg.UDPEnabled, cfg.TCPEnabled, cfg.TrustProxyHeaders, cfg.TrustedProxyCIDRs)
136 + registry, err := newLeaseRegistry(cfg.UDPEnabled, cfg.TCPEnabled, cfg.MinPort, cfg.MaxPort, identity.Name, cfg.SNIPort, identity.PrivateKey, identity.PublicKey, identity.Address, cfg.PortalURL, cfg.TrustProxyHeaders, cfg.TrustedProxyCIDRs)
137 if err != nil {
138 return nil, err
139 }
@@ -151,15 +147,15 @@ func NewServer(cfg ServerConfig) (*Server, error) {
147 relaySet = discovery.NewRelaySet(cfg.Bootstraps)
148 }
149
154 - return &Server{
150 + server := &Server{
151 cfg: cfg,
152 identity: identity,
153 registry: registry,
154 relaySet: relaySet,
155 announceLimiter: discovery.NewAnnounceLimiter(0, 0),
160 - udpPorts: transport.NewPortAllocator(cfg.MinPort, cfg.MaxPort, 5*time.Minute),
161 - tcpPorts: transport.NewPortAllocator(cfg.MinPort, cfg.MaxPort, 5*time.Minute),
162 - }, nil
156 + }
157 + server.registry.proxy = &server.proxy
158 + return server, nil
159 }
160
161 func (s *Server) Start(ctx context.Context, apiMux *http.ServeMux) error {
@@ -349,7 +345,7 @@ func (s *Server) Shutdown(ctx context.Context) error {
345 }
346
347 for _, lease := range s.registry.CloseAll() {
352 - s.cleanupRemovedRecord(ctx, lease, "delete lease remote state during shutdown")
348 + s.deleteENSGaslessHostname(ctx, lease, "delete lease ens gasless hostname during shutdown")
349 }
350
351 if s.quicBackhaul != nil {
@@ -414,6 +410,14 @@ func (s *Server) prepareAPITLS(ctx context.Context) (keyless.TLSMaterialConfig,
410 return apiTLS, manager, nil
411 }
412
413 +func (s *Server) runAPIServer() error {
414 + err := s.apiServer.Serve(s.apiListener)
415 + if err == nil || errors.Is(err, http.ErrServerClosed) || errors.Is(err, net.ErrClosed) {
416 + return nil
417 + }
418 + return err
419 +}
420 +
421 func (s *Server) runPublicIngress(ctx context.Context) error {
422 for {
423 conn, err := s.sniListener.Accept()
@@ -486,8 +490,10 @@ func (s *Server) runOverlayIngress(ctx context.Context) error {
490 return err
491 }
492 go func(stream overlay.HopStream) {
489 - record, ok := s.registry.RecordByHopToken(stream.Token, time.Now())
490 - if !ok {
493 + s.registry.mu.RLock()
494 + record := s.registry.recordByHopToken(stream.Token, time.Now())
495 + s.registry.mu.RUnlock()
496 + if record == nil {
497 log.Warn().Str("remote_addr", stream.RemoteAddr).Msg("hop stream rejected")
498 _ = stream.Conn.Close()
499 return
@@ -571,7 +577,7 @@ func (s *Server) runRegistryJanitor(ctx context.Context, interval time.Duration)
577 return nil
578 case <-ticker.C:
579 for _, lease := range s.registry.cleanupExpired(time.Now()) {
574 - s.cleanupRemovedRecord(context.Background(), lease, "delete expired lease remote state")
580 + s.deleteENSGaslessHostname(context.Background(), lease, "delete expired lease ens gasless hostname")
581 }
582 }
583 }
@@ -604,6 +610,45 @@ func (s *Server) runQUICBackhaulListener() error {
610 }
611 }
612
613 +func (s *Server) handleQUICBackhaulConn(conn *quic.Conn) {
614 + control, err := transport.AcceptQUICBackhaulControl(context.Background(), conn)
615 + if err != nil {
616 + _ = conn.CloseWithError(1, "control read failed")
617 + return
618 + }
619 +
620 + lease, err := s.registry.admitLeaseByToken(control.AccessToken, true)
621 + if err != nil {
622 + code, reason := types.APIErrorCodeInvalidRequest, "invalid control message"
623 + switch {
624 + case errors.Is(err, errLeaseNotFound):
625 + code, reason = types.APIErrorCodeLeaseNotFound, "lease not found"
626 + case errors.Is(err, errLeaseRejected):
627 + code, reason = types.APIErrorCodeLeaseRejected, "lease rejected"
628 + case errors.Is(err, errUnauthorized):
629 + code, reason = types.APIErrorCodeUnauthorized, "unauthorized"
630 + case errors.Is(err, errTransportMismatch):
631 + code, reason = types.APIErrorCodeTransportMismatch, "transport mismatch"
632 + }
633 + _ = control.Reject(code, reason)
634 + return
635 + }
636 +
637 + if err := lease.datagram.BindBackhaul(conn); err != nil {
638 + _ = control.Reject("broker_closed", "broker closed")
639 + return
640 + }
641 +
642 + _ = control.Accept()
643 + s.registry.Touch(lease.Key(), conn.RemoteAddr().String(), time.Now())
644 + log.Info().
645 + Str("component", "quic-backhaul-listener").
646 + Str("address", lease.Address).
647 + Str("lease_name", lease.Name).
648 + Str("remote_addr", conn.RemoteAddr().String()).
649 + Msg("quic backhaul connected")
650 +}
651 +
652 func (s *Server) startOverlay() (*overlay.Overlay, error) {
653 peerMux := http.NewServeMux()
654 peerMux.HandleFunc(types.PathRoot, s.handleRoot)
@@ -640,7 +685,7 @@ func (s *Server) runRelayDiscoveryLoop(ctx context.Context) error {
685
686 for {
687 now := time.Now().UTC()
643 - self, err := s.signedRelayDescriptor(now)
688 + self, err := s.newSelfDescriptor(now)
689 if err != nil {
690 return fmt.Errorf("build relay discovery descriptor: %w", err)
691 }
@@ -661,3 +706,59 @@ func (s *Server) runRelayDiscoveryLoop(ctx context.Context) error {
706 }
707 }
708 }
709 +
710 +func (s *Server) newSelfDescriptor(now time.Time) (types.RelayDescriptor, error) {
711 + if now.IsZero() {
712 + now = time.Now().UTC()
713 + } else {
714 + now = now.UTC()
715 + }
716 +
717 + var wireGuardPublicKey string
718 + var wireGuardPort int
719 + if s.overlay != nil {
720 + cfg := s.overlay.Config()
721 + wireGuardPublicKey = cfg.PublicKey
722 + wireGuardPort = cfg.ListenPort
723 + }
724 +
725 + return auth.SignRelayDescriptor(types.RelayDescriptor{
726 + Address: s.identity.Address,
727 + Version: types.DiscoveryVersion,
728 + IssuedAt: now,
729 + ExpiresAt: now.Add(discovery.DiscoveryDescriptorTTL),
730 + APIHTTPSAddr: s.cfg.PortalURL,
731 + WireGuardPublicKey: wireGuardPublicKey,
732 + WireGuardPort: wireGuardPort,
733 + SupportsOverlay: s.overlay != nil,
734 + SupportsUDP: s.cfg.UDPEnabled && s.quicBackhaul != nil,
735 + SupportsTCP: s.cfg.TCPEnabled,
736 + ActiveConnections: s.proxy.activeConnectionCount(),
737 + TCPBPS: s.proxy.currentTCPBPS(now),
738 + }, s.identity.PrivateKey)
739 +}
740 +
741 +func (s *Server) syncENSGaslessHostname(ctx context.Context, record *leaseRecord) error {
742 + if record == nil || !record.isPublicEntry() || s.acmeManager == nil {
743 + return nil
744 + }
745 + syncCtx, cancel := context.WithTimeout(ctx, defaultClaimTimeout)
746 + defer cancel()
747 + return s.acmeManager.SyncENSGaslessHostname(syncCtx, record.Hostname, record.Address)
748 +}
749 +
750 +func (s *Server) deleteENSGaslessHostname(ctx context.Context, record *leaseRecord, logMessage string) {
751 + if record == nil || !record.isPublicEntry() || s.acmeManager == nil {
752 + return
753 + }
754 + deleteCtx, cancel := context.WithTimeout(ctx, defaultClaimTimeout)
755 + err := s.acmeManager.DeleteENSGaslessHostname(deleteCtx, record.Hostname)
756 + cancel()
757 + if err != nil {
758 + log.Warn().
759 + Err(err).
760 + Str("hostname", record.Hostname).
761 + Str("address", record.Address).
762 + Msg(logMessage)
763 + }
764 +}
portal/server_test.go
+9 -21
@@ -228,7 +228,7 @@ func TestRegisterLeaseOmitsSNIPortWithoutUDP(t *testing.T) {
228 t.Fatalf("NewServer() error = %v", err)
229 }
230
231 - resp, err := server.registerLease(types.RegisterChallengeRequest{
231 + record, resp, err := server.registry.Register(types.RegisterChallengeRequest{
232 Identity: types.Identity{
233 Name: "demo-tcp",
234 Address: server.identity.Address,
@@ -236,12 +236,10 @@ func TestRegisterLeaseOmitsSNIPortWithoutUDP(t *testing.T) {
236 TCPEnabled: true,
237 }, "203.0.113.10", "")
238 if err != nil {
239 - t.Fatalf("registerLease() error = %v", err)
239 + t.Fatalf("registry.Register() error = %v", err)
240 }
241 t.Cleanup(func() {
242 - if record, ok := server.registry.RecordByKey(resp.Identity.Key(), time.Now()); ok {
243 - record.Close()
244 - }
242 + record.Close()
243 })
244
245 if resp.SNIPort != 0 {
@@ -326,25 +324,21 @@ func TestRegisterLeaseDerivesFixedHostnameFromName(t *testing.T) {
324 t.Fatalf("NewServer() error = %v", err)
325 }
326
329 - resp, err := server.registerLease(types.RegisterChallengeRequest{
327 + record, resp, err := server.registry.Register(types.RegisterChallengeRequest{
328 Identity: types.Identity{
329 Name: "Demo-App",
330 Address: server.identity.Address,
331 },
332 }, "203.0.113.10", "")
333 if err != nil {
336 - t.Fatalf("registerLease() error = %v", err)
334 + t.Fatalf("registry.Register() error = %v", err)
335 }
336
337 wantHostname := "demo-app.portal.example.com"
338 if resp.Hostname != wantHostname {
341 - t.Fatalf("registerLease() hostname = %q, want %q", resp.Hostname, wantHostname)
339 + t.Fatalf("registry.Register() hostname = %q, want %q", resp.Hostname, wantHostname)
340 }
341
344 - record, ok := server.registry.RecordByKey(resp.Identity.Key(), time.Now())
345 - if !ok {
346 - t.Fatal("registry.RecordByKey() = false, want registered lease")
347 - }
342 lease := server.registry.publicLease(record)
343 if lease.Name != "demo-app" {
344 t.Fatalf("publicLease().Name = %q, want %q", lease.Name, "demo-app")
@@ -369,7 +363,7 @@ func TestRegisterLeaseBuildsUDPEnabledRuntime(t *testing.T) {
363 }
364 server.registry.policy.SetUDPPolicy(true, 0)
365
372 - resp, err := server.registerLease(types.RegisterChallengeRequest{
366 + record, resp, err := server.registry.Register(types.RegisterChallengeRequest{
367 Identity: types.Identity{
368 Name: "demo-udp",
369 Address: server.identity.Address,
@@ -377,18 +371,12 @@ func TestRegisterLeaseBuildsUDPEnabledRuntime(t *testing.T) {
371 UDPEnabled: true,
372 }, "203.0.113.10", "")
373 if err != nil {
380 - t.Fatalf("registerLease() error = %v", err)
374 + t.Fatalf("registry.Register() error = %v", err)
375 }
376 t.Cleanup(func() {
383 - if record, ok := server.registry.RecordByKey(resp.Identity.Key(), time.Now()); ok {
384 - record.Close()
385 - }
377 + record.Close()
378 })
379
388 - record, ok := server.registry.RecordByKey(resp.Identity.Key(), time.Now())
389 - if !ok {
390 - t.Fatal("registry.RecordByKey() = false, want registered lease")
391 - }
380 if record.stream == nil {
381 t.Fatal("stream = nil, want stream runtime")
382 }