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
}