lease: remove discovery

rabbitprincess committed Mar 29, 2026 at 21:52 UTC 98292f149dcbe618adbb2e11a2c4e41bcc353cc0
12 files changed +1099 -877
portal/api_server.go
+4 -83
@@ -117,7 +117,10 @@ func (s *Server) discover(_ context.Context, req types.DiscoverRequest) (types.D
117 ProtocolVersion: 1,
118 GeneratedAt: time.Now().UTC(),
119 Self: self,
120 - Peers: s.discoveryAdvertisedPeerDescriptors(nil),
120 + Peers: nil,
121 + }
122 + if s.peerRegistry != nil {
123 + resp.Peers = s.peerRegistry.advertisedPeers()
124 }
125 if req.RootHost != "" && req.RootHost != s.rootHost {
126 return resp, nil
@@ -140,7 +143,6 @@ func (s *Server) discover(_ context.Context, req types.DiscoverRequest) (types.D
143 record, found := s.registry.leaseByID[leaseID]
144 if found && record != nil && !now.After(record.ExpiresAt) && !record.Metadata.Hide && s.registry.policy.IsLeaseRoutable(record.ID) {
145 lease = record.Lease
143 - lease.Bootstraps = append([]string(nil), lease.Bootstraps...)
146 lease.Metadata = lease.Metadata.Copy()
147 ok = true
148 }
@@ -156,7 +158,6 @@ func (s *Server) discover(_ context.Context, req types.DiscoverRequest) (types.D
158 ownerAddress = s.ownerIdentity.Address
159 }
160
159 - resp.Peers = s.discoveryAdvertisedPeerDescriptors(lease.Bootstraps)
161 resp.Service = &types.DiscoveredService{
162 Found: true,
163 Name: lease.Name,
@@ -205,38 +206,6 @@ func (s *Server) discoverySelfDescriptor() (types.RelayDescriptor, error) {
206 return discovery.SignedDescriptor(descriptor, s.ownerIdentity.PrivateKey)
207 }
208
208 -func (s *Server) discoveryAdvertisedPeerDescriptors(urls []string) []types.RelayDescriptor {
209 - advertised := s.discoveryCache.AdvertisedDescriptors()
210 - if len(advertised) == 0 {
211 - return nil
212 - }
213 - if len(urls) == 0 {
214 - return advertised
215 - }
216 -
217 - allowed := make(map[string]struct{}, len(urls))
218 - for _, relayURL := range urls {
219 - relayURL = strings.TrimSpace(relayURL)
220 - if relayURL == "" {
221 - continue
222 - }
223 - normalized, err := utils.NormalizeRelayURL(relayURL)
224 - if err != nil {
225 - continue
226 - }
227 - allowed[normalized] = struct{}{}
228 - }
229 -
230 - records := make([]types.RelayDescriptor, 0, len(advertised))
231 - for _, descriptor := range advertised {
232 - if _, ok := allowed[descriptor.APIHTTPSAddr]; !ok {
233 - continue
234 - }
235 - records = append(records, descriptor)
236 - }
237 - return records
238 -}
239 -
209 func (s *Server) handleDomain(w http.ResponseWriter, r *http.Request) {
210 w.Header().Set("Access-Control-Allow-Origin", "*")
211
@@ -540,10 +509,6 @@ func (s *Server) registerLease(req types.RegisterRequest, clientIP string) (type
509 if req.TTL > 0 {
510 ttl = time.Duration(req.TTL) * time.Second
511 }
543 - bootstraps, err := utils.NormalizeRelayURLs(req.Bootstraps...)
544 - if err != nil {
545 - return types.RegisterResponse{}, fmt.Errorf("normalize bootstraps: %w", err)
546 - }
512 ownerAddress := strings.TrimSpace(req.OwnerAddress)
513 if ownerAddress != "" {
514 ownerAddress, err = utils.NormalizeEVMAddress(ownerAddress)
@@ -572,7 +537,6 @@ func (s *Server) registerLease(req types.RegisterRequest, clientIP string) (type
537 ID: leaseID,
538 Name: name,
539 Hostname: hostname,
575 - Bootstraps: append([]string(nil), bootstraps...),
540 Metadata: req.Metadata,
541 OwnerAddress: ownerAddress,
542 ExpiresAt: expiresAt,
@@ -606,57 +570,14 @@ func (s *Server) registerLease(req types.RegisterRequest, clientIP string) (type
570 record.Close()
571 return types.RegisterResponse{}, err
572 }
609 - if s.DiscoveryEnabled() {
610 - if _, err := s.discoveryCache.UpsertSeedURLs(bootstraps); err != nil {
611 - record.Close()
612 - _, _ = s.registry.Unregister(record.ID, record.ReverseToken)
613 - return types.RegisterResponse{}, err
614 - }
615 - }
616 -
617 - responseBootstraps, err := utils.NormalizeRelayURLs(s.cfg.PortalURL)
618 - if err != nil {
619 - record.Close()
620 - _, _ = s.registry.Unregister(record.ID, record.ReverseToken)
621 - return types.RegisterResponse{}, err
622 - }
623 - if s.DiscoveryEnabled() {
624 - advertisedURLs := make([]string, 0, len(s.discoveryCache.AdvertisedDescriptors()))
625 - for _, descriptor := range s.discoveryCache.AdvertisedDescriptors() {
626 - if apiURL := strings.TrimSpace(descriptor.APIHTTPSAddr); apiURL != "" {
627 - advertisedURLs = append(advertisedURLs, apiURL)
628 - }
629 - }
630 - advertisedURLs, err = utils.ExcludeLocalRelayURLs(advertisedURLs...)
631 - if err != nil {
632 - record.Close()
633 - _, _ = s.registry.Unregister(record.ID, record.ReverseToken)
634 - return types.RegisterResponse{}, err
635 - }
636 - responseBootstraps, err = utils.NormalizeRelayURLs(append(responseBootstraps, advertisedURLs...)...)
637 - if err != nil {
638 - record.Close()
639 - _, _ = s.registry.Unregister(record.ID, record.ReverseToken)
640 - return types.RegisterResponse{}, err
641 - }
642 - } else {
643 - responseBootstraps, err = utils.NormalizeRelayURLs(append(responseBootstraps, append(s.cfg.Bootstraps, record.Bootstraps...)...)...)
644 - }
645 - if err != nil {
646 - record.Close()
647 - _, _ = s.registry.Unregister(record.ID, record.ReverseToken)
648 - return types.RegisterResponse{}, err
649 - }
573
574 resp := types.RegisterResponse{
575 LeaseID: leaseID,
576 Hostname: hostname,
577 Metadata: record.Metadata,
578 ExpiresAt: expiresAt,
656 - Bootstraps: responseBootstraps,
579 UDPEnabled: record.UDPEnabled,
580 }
659 - resp.ConnectURL = strings.TrimRight(s.cfg.PortalURL, "/") + types.PathSDKConnect
581 if record.datagram != nil {
582 resp.UDPAddr = fmt.Sprintf("%s:%d", s.rootHost, record.datagram.UDPPort())
583 }
portal/discovery/store.go deleted
-293
@@ -1,293 +0,0 @@
1 -package discovery
2 -
3 -import (
4 - "errors"
5 - "reflect"
6 - "sort"
7 - "strings"
8 - "sync"
9 - "time"
10 -
11 - "github.com/gosuda/portal/v2/types"
12 - "github.com/gosuda/portal/v2/utils"
13 -)
14 -
15 -type peerRecord struct {
16 - seedURL string
17 - pinnedSignerPublicKey string
18 - state types.PeerState
19 -}
20 -
21 -// Cache keeps only verified relay descriptors and bootstrap hints.
22 -// It is not a source of truth; trust comes from seed URLs plus signer pinning.
23 -type Cache struct {
24 - mu sync.RWMutex
25 - peers map[string]peerRecord
26 -}
27 -
28 -func (s *Cache) Lookup(relayID string) (types.PeerState, bool, bool) {
29 - if strings.TrimSpace(relayID) == "" {
30 - return types.PeerState{}, false, false
31 - }
32 -
33 - s.mu.RLock()
34 - defer s.mu.RUnlock()
35 -
36 - record, ok := s.peers[relayID]
37 - if !ok {
38 - return types.PeerState{}, false, false
39 - }
40 - return record.state, strings.TrimSpace(record.pinnedSignerPublicKey) != "", true
41 -}
42 -
43 -func NewCache() *Cache {
44 - return &Cache{
45 - peers: make(map[string]peerRecord),
46 - }
47 -}
48 -
49 -func SeedDescriptor(apiURL string) (types.RelayDescriptor, error) {
50 - normalized, err := utils.NormalizeRelayURL(apiURL)
51 - if err != nil {
52 - return types.RelayDescriptor{}, err
53 - }
54 - return types.RelayDescriptor{
55 - RelayID: normalized,
56 - APIHTTPSAddr: normalized,
57 - Version: 1,
58 - }, nil
59 -}
60 -
61 -func (s *Cache) UpsertSeedURLs(inputs []string) ([]string, error) {
62 - if len(inputs) == 0 {
63 - return nil, nil
64 - }
65 -
66 - normalized, err := utils.NormalizeRelayURLs(inputs...)
67 - if err != nil {
68 - return nil, err
69 - }
70 - normalized, err = utils.ExcludeLocalRelayURLs(normalized...)
71 - if err != nil {
72 - return nil, err
73 - }
74 -
75 - now := time.Now().UTC()
76 - added := make([]string, 0, len(normalized))
77 -
78 - s.mu.Lock()
79 - defer s.mu.Unlock()
80 -
81 - for _, apiURL := range normalized {
82 - descriptor, err := SeedDescriptor(apiURL)
83 - if err != nil {
84 - return nil, err
85 - }
86 -
87 - record, ok := s.peers[descriptor.RelayID]
88 - if !ok {
89 - s.peers[descriptor.RelayID] = peerRecord{
90 - seedURL: descriptor.APIHTTPSAddr,
91 - state: types.PeerState{
92 - Descriptor: descriptor,
93 - State: types.PeerStateKnown,
94 - FirstSeenAt: now,
95 - LastSeenAt: now,
96 - },
97 - }
98 - added = append(added, descriptor.APIHTTPSAddr)
99 - continue
100 - }
101 -
102 - if strings.TrimSpace(record.seedURL) == "" {
103 - record.seedURL = descriptor.APIHTTPSAddr
104 - }
105 - if strings.TrimSpace(record.state.Descriptor.APIHTTPSAddr) == "" {
106 - record.state.Descriptor.APIHTTPSAddr = descriptor.APIHTTPSAddr
107 - }
108 - if strings.TrimSpace(record.state.Descriptor.RelayID) == "" {
109 - record.state.Descriptor.RelayID = descriptor.RelayID
110 - }
111 - record.state.LastSeenAt = now
112 - s.peers[descriptor.RelayID] = record
113 - }
114 -
115 - return added, nil
116 -}
117 -
118 -func (s *Cache) Snapshot() map[string]types.PeerState {
119 - s.mu.RLock()
120 - defer s.mu.RUnlock()
121 -
122 - out := make(map[string]types.PeerState, len(s.peers))
123 - for relayID, record := range s.peers {
124 - out[relayID] = record.state
125 - }
126 - return out
127 -}
128 -
129 -func (s *Cache) KnownDescriptors() []types.RelayDescriptor {
130 - s.mu.RLock()
131 - defer s.mu.RUnlock()
132 -
133 - out := make([]types.RelayDescriptor, 0, len(s.peers))
134 - for _, record := range s.peers {
135 - if strings.TrimSpace(record.state.Descriptor.APIHTTPSAddr) == "" {
136 - continue
137 - }
138 - out = append(out, record.state.Descriptor)
139 - }
140 - sort.Slice(out, func(i, j int) bool {
141 - return out[i].APIHTTPSAddr < out[j].APIHTTPSAddr
142 - })
143 - return out
144 -}
145 -
146 -func (s *Cache) AdvertisedDescriptors() []types.RelayDescriptor {
147 - s.mu.RLock()
148 - defer s.mu.RUnlock()
149 -
150 - out := make([]types.RelayDescriptor, 0, len(s.peers))
151 - for _, record := range s.peers {
152 - if record.state.State != types.PeerStateAdvertised {
153 - continue
154 - }
155 - out = append(out, record.state.Descriptor)
156 - }
157 - sort.Slice(out, func(i, j int) bool {
158 - return out[i].APIHTTPSAddr < out[j].APIHTTPSAddr
159 - })
160 - return out
161 -}
162 -
163 -func (s *Cache) PinIdentity(relayID, seedURL string, desc types.RelayDescriptor) error {
164 - relayID = strings.TrimSpace(relayID)
165 - if relayID == "" {
166 - return errors.New("relay id is required")
167 - }
168 - normalizedSeedURL, err := utils.NormalizeRelayURL(seedURL)
169 - if err != nil {
170 - return err
171 - }
172 - if strings.TrimSpace(desc.APIHTTPSAddr) == "" {
173 - return errors.New("descriptor api_https_addr is required")
174 - }
175 - if desc.APIHTTPSAddr != normalizedSeedURL {
176 - return errors.New("descriptor api_https_addr does not match seed url")
177 - }
178 - if strings.TrimSpace(desc.SignerPublicKey) == "" {
179 - return errors.New("descriptor signer_public_key is required")
180 - }
181 -
182 - now := time.Now().UTC()
183 -
184 - s.mu.Lock()
185 - defer s.mu.Unlock()
186 -
187 - record, ok := s.peers[relayID]
188 - if !ok {
189 - record = peerRecord{
190 - state: types.PeerState{
191 - FirstSeenAt: now,
192 - },
193 - }
194 - }
195 - if record.seedURL != "" && record.seedURL != normalizedSeedURL {
196 - return errors.New("seed url does not match cached relay url")
197 - }
198 - if record.pinnedSignerPublicKey != "" && record.pinnedSignerPublicKey != desc.SignerPublicKey {
199 - return errors.New("descriptor signer_public_key does not match pinned signer")
200 - }
201 -
202 - record.seedURL = normalizedSeedURL
203 - record.pinnedSignerPublicKey = desc.SignerPublicKey
204 - s.peers[relayID] = record
205 - return nil
206 -}
207 -
208 -func (s *Cache) RecordVerified(desc types.RelayDescriptor, advertise bool) (bool, bool, error) {
209 - if strings.TrimSpace(desc.RelayID) == "" {
210 - return false, false, errors.New("relay id is required")
211 - }
212 -
213 - now := time.Now().UTC()
214 -
215 - s.mu.Lock()
216 - defer s.mu.Unlock()
217 -
218 - record, ok := s.peers[desc.RelayID]
219 - added := !ok
220 - if !ok {
221 - record = peerRecord{
222 - state: types.PeerState{
223 - FirstSeenAt: now,
224 - },
225 - }
226 - }
227 -
228 - if strings.TrimSpace(record.seedURL) == "" {
229 - record.seedURL = strings.TrimSpace(desc.APIHTTPSAddr)
230 - } else if strings.TrimSpace(desc.APIHTTPSAddr) != "" && record.seedURL != strings.TrimSpace(desc.APIHTTPSAddr) {
231 - return false, false, errors.New("descriptor api_https_addr does not match cached seed url")
232 - }
233 - if record.pinnedSignerPublicKey != "" && record.pinnedSignerPublicKey != desc.SignerPublicKey {
234 - return false, false, errors.New("descriptor signer_public_key does not match pinned signer")
235 - }
236 -
237 - previousState := record.state.State
238 - previousDescriptor := record.state.Descriptor
239 - switch {
240 - case advertise:
241 - record.state.State = types.PeerStateAdvertised
242 - case record.state.State == types.PeerStateAdvertised:
243 - record.state.State = types.PeerStateAdvertised
244 - default:
245 - record.state.State = types.PeerStateVerified
246 - }
247 -
248 - record.state.Descriptor = desc
249 - record.state.LastSeenAt = now
250 - record.state.ConsecutiveFailures = 0
251 - s.peers[desc.RelayID] = record
252 -
253 - changed := added ||
254 - previousState != record.state.State ||
255 - !reflect.DeepEqual(previousDescriptor, record.state.Descriptor)
256 - return added, changed, nil
257 -}
258 -
259 -func (s *Cache) RecordFailure(relayID string) {
260 - if strings.TrimSpace(relayID) == "" {
261 - return
262 - }
263 -
264 - s.mu.Lock()
265 - defer s.mu.Unlock()
266 -
267 - record, ok := s.peers[relayID]
268 - if !ok {
269 - return
270 - }
271 - record.state.ConsecutiveFailures++
272 - s.peers[relayID] = record
273 -}
274 -
275 -func (s *Cache) Expire(relayID string) bool {
276 - if strings.TrimSpace(relayID) == "" {
277 - return false
278 - }
279 -
280 - s.mu.Lock()
281 - defer s.mu.Unlock()
282 -
283 - record, ok := s.peers[relayID]
284 - if !ok {
285 - return false
286 - }
287 - if record.state.State == types.PeerStateExpired {
288 - return false
289 - }
290 - record.state.State = types.PeerStateExpired
291 - s.peers[relayID] = record
292 - return true
293 -}
portal/lease.go
-1
@@ -234,7 +234,6 @@ func (r *leaseRegistry) Snapshot(record *leaseRecord) types.Lease {
234 }
235
236 snapshot := record.Lease
237 - snapshot.Bootstraps = append([]string(nil), snapshot.Bootstraps...)
237 snapshot.Metadata = snapshot.Metadata.Copy()
238 clientIP := record.ClientIP
239 snapshot.BPS = r.policy.BPSManager().LeaseBPS(record.ID)
portal/peer.go new
+407
@@ -0,0 +1,407 @@
1 +package portal
2 +
3 +import (
4 + "errors"
5 + "net/http"
6 + "reflect"
7 + "sort"
8 + "strings"
9 + "sync"
10 + "time"
11 +
12 + "github.com/gosuda/portal/v2/types"
13 + "github.com/gosuda/portal/v2/utils"
14 +)
15 +
16 +type peerRecord struct {
17 + bootstrap bool
18 + seedURL string
19 + pinnedSignerPublicKey string
20 + state types.PeerState
21 +}
22 +
23 +// peerRegistry keeps only verified relay descriptors and bootstrap hints.
24 +// It is not a source of truth; trust comes from seed URLs plus signer pinning.
25 +type peerRegistry struct {
26 + mu sync.RWMutex
27 + peers map[string]peerRecord
28 +}
29 +
30 +type peerRegistrationResult struct {
31 + updated bool
32 + addedHintCount int
33 + peerSetChanged bool
34 +}
35 +
36 +type peerFailureOutcome struct {
37 + expired bool
38 + expireReason string
39 + consecutiveFailures int
40 +}
41 +
42 +func newPeerRegistry() *peerRegistry {
43 + return &peerRegistry{
44 + peers: make(map[string]peerRecord),
45 + }
46 +}
47 +
48 +func (r *peerRegistry) lookup(relayID string) (types.PeerState, bool, bool) {
49 + if strings.TrimSpace(relayID) == "" {
50 + return types.PeerState{}, false, false
51 + }
52 +
53 + r.mu.RLock()
54 + defer r.mu.RUnlock()
55 +
56 + record, ok := r.peers[relayID]
57 + if !ok {
58 + return types.PeerState{}, false, false
59 + }
60 + return record.state, strings.TrimSpace(record.pinnedSignerPublicKey) != "", true
61 +}
62 +
63 +func seedDescriptor(apiURL string) (types.RelayDescriptor, error) {
64 + normalized, err := utils.NormalizeRelayURL(apiURL)
65 + if err != nil {
66 + return types.RelayDescriptor{}, err
67 + }
68 + return types.RelayDescriptor{
69 + RelayID: normalized,
70 + APIHTTPSAddr: normalized,
71 + Version: 1,
72 + }, nil
73 +}
74 +
75 +func requireOverlayPeerDescriptor(desc types.RelayDescriptor) error {
76 + if !desc.SupportsOverlayPeer {
77 + return errors.New("descriptor does not support overlay peer")
78 + }
79 + if strings.TrimSpace(desc.WireGuardPublicKey) == "" {
80 + return errors.New("descriptor wireguard public key is required")
81 + }
82 + if strings.TrimSpace(desc.WireGuardEndpoint) == "" {
83 + return errors.New("descriptor wireguard endpoint is required")
84 + }
85 + if strings.TrimSpace(desc.OverlayIPv4) == "" {
86 + return errors.New("descriptor overlay ipv4 is required")
87 + }
88 + return nil
89 +}
90 +
91 +func (r *peerRegistry) registerBootstrapURLs(inputs []string) ([]string, error) {
92 + if len(inputs) == 0 {
93 + return nil, nil
94 + }
95 +
96 + normalized, err := utils.NormalizeRelayURLs(inputs...)
97 + if err != nil {
98 + return nil, err
99 + }
100 + normalized, err = utils.ExcludeLocalRelayURLs(normalized...)
101 + if err != nil {
102 + return nil, err
103 + }
104 +
105 + now := time.Now().UTC()
106 + added := make([]string, 0, len(normalized))
107 +
108 + r.mu.Lock()
109 + defer r.mu.Unlock()
110 +
111 + for _, apiURL := range normalized {
112 + descriptor, err := seedDescriptor(apiURL)
113 + if err != nil {
114 + return nil, err
115 + }
116 +
117 + record, ok := r.peers[descriptor.RelayID]
118 + if !ok {
119 + r.peers[descriptor.RelayID] = peerRecord{
120 + bootstrap: true,
121 + seedURL: descriptor.APIHTTPSAddr,
122 + state: types.PeerState{
123 + Descriptor: descriptor,
124 + State: types.PeerStateKnown,
125 + FirstSeenAt: now,
126 + LastSeenAt: now,
127 + },
128 + }
129 + added = append(added, descriptor.APIHTTPSAddr)
130 + continue
131 + }
132 +
133 + record.bootstrap = true
134 + if strings.TrimSpace(record.seedURL) == "" {
135 + record.seedURL = descriptor.APIHTTPSAddr
136 + }
137 + if strings.TrimSpace(record.state.Descriptor.APIHTTPSAddr) == "" {
138 + record.state.Descriptor.APIHTTPSAddr = descriptor.APIHTTPSAddr
139 + }
140 + if strings.TrimSpace(record.state.Descriptor.RelayID) == "" {
141 + record.state.Descriptor.RelayID = descriptor.RelayID
142 + }
143 + record.state.LastSeenAt = now
144 + r.peers[descriptor.RelayID] = record
145 + }
146 +
147 + return added, nil
148 +}
149 +
150 +func (r *peerRegistry) snapshot() map[string]types.PeerState {
151 + r.mu.RLock()
152 + defer r.mu.RUnlock()
153 +
154 + out := make(map[string]types.PeerState, len(r.peers))
155 + for relayID, record := range r.peers {
156 + out[relayID] = record.state
157 + }
158 + return out
159 +}
160 +
161 +func (r *peerRegistry) bootstrapPeers() []types.RelayDescriptor {
162 + r.mu.RLock()
163 + defer r.mu.RUnlock()
164 +
165 + out := make([]types.RelayDescriptor, 0, len(r.peers))
166 + for _, record := range r.peers {
167 + if !record.bootstrap || strings.TrimSpace(record.state.Descriptor.APIHTTPSAddr) == "" {
168 + continue
169 + }
170 + out = append(out, record.state.Descriptor)
171 + }
172 + sort.Slice(out, func(i, j int) bool {
173 + return out[i].APIHTTPSAddr < out[j].APIHTTPSAddr
174 + })
175 + return out
176 +}
177 +
178 +func (r *peerRegistry) advertisedPeers() []types.RelayDescriptor {
179 + r.mu.RLock()
180 + defer r.mu.RUnlock()
181 +
182 + out := make([]types.RelayDescriptor, 0, len(r.peers))
183 + for _, record := range r.peers {
184 + if record.state.State != types.PeerStateAdvertised {
185 + continue
186 + }
187 + out = append(out, record.state.Descriptor)
188 + }
189 + sort.Slice(out, func(i, j int) bool {
190 + return out[i].APIHTTPSAddr < out[j].APIHTTPSAddr
191 + })
192 + return out
193 +}
194 +
195 +func (r *peerRegistry) syncablePeers() []types.RelayDescriptor {
196 + r.mu.RLock()
197 + defer r.mu.RUnlock()
198 +
199 + out := make([]types.RelayDescriptor, 0, len(r.peers))
200 + for _, record := range r.peers {
201 + if record.bootstrap ||
202 + record.state.State == types.PeerStateExpired ||
203 + !record.state.Descriptor.SupportsOverlayPeer {
204 + continue
205 + }
206 + out = append(out, record.state.Descriptor)
207 + }
208 + sort.Slice(out, func(i, j int) bool {
209 + return out[i].APIHTTPSAddr < out[j].APIHTTPSAddr
210 + })
211 + return out
212 +}
213 +
214 +func (r *peerRegistry) pin(relayID, seedURL string, desc types.RelayDescriptor) error {
215 + relayID = strings.TrimSpace(relayID)
216 + if relayID == "" {
217 + return errors.New("relay id is required")
218 + }
219 + normalizedSeedURL, err := utils.NormalizeRelayURL(seedURL)
220 + if err != nil {
221 + return err
222 + }
223 + if strings.TrimSpace(desc.APIHTTPSAddr) == "" {
224 + return errors.New("descriptor api_https_addr is required")
225 + }
226 + if desc.APIHTTPSAddr != normalizedSeedURL {
227 + return errors.New("descriptor api_https_addr does not match seed url")
228 + }
229 + if strings.TrimSpace(desc.SignerPublicKey) == "" {
230 + return errors.New("descriptor signer_public_key is required")
231 + }
232 +
233 + now := time.Now().UTC()
234 +
235 + r.mu.Lock()
236 + defer r.mu.Unlock()
237 +
238 + record, ok := r.peers[relayID]
239 + if !ok {
240 + record = peerRecord{
241 + state: types.PeerState{
242 + FirstSeenAt: now,
243 + },
244 + }
245 + }
246 + if record.seedURL != "" && record.seedURL != normalizedSeedURL {
247 + return errors.New("seed url does not match cached relay url")
248 + }
249 + if record.pinnedSignerPublicKey != "" && record.pinnedSignerPublicKey != desc.SignerPublicKey {
250 + return errors.New("descriptor signer_public_key does not match pinned signer")
251 + }
252 +
253 + record.seedURL = normalizedSeedURL
254 + record.pinnedSignerPublicKey = desc.SignerPublicKey
255 + r.peers[relayID] = record
256 + return nil
257 +}
258 +
259 +func (r *peerRegistry) register(desc types.RelayDescriptor, advertise bool) (bool, bool, error) {
260 + if strings.TrimSpace(desc.RelayID) == "" {
261 + return false, false, errors.New("relay id is required")
262 + }
263 +
264 + now := time.Now().UTC()
265 +
266 + r.mu.Lock()
267 + defer r.mu.Unlock()
268 +
269 + record, ok := r.peers[desc.RelayID]
270 + added := !ok
271 + if !ok {
272 + record = peerRecord{
273 + state: types.PeerState{
274 + FirstSeenAt: now,
275 + },
276 + }
277 + }
278 +
279 + if strings.TrimSpace(record.seedURL) == "" {
280 + record.seedURL = strings.TrimSpace(desc.APIHTTPSAddr)
281 + } else if strings.TrimSpace(desc.APIHTTPSAddr) != "" && record.seedURL != strings.TrimSpace(desc.APIHTTPSAddr) {
282 + return false, false, errors.New("descriptor api_https_addr does not match cached seed url")
283 + }
284 + if record.pinnedSignerPublicKey != "" && record.pinnedSignerPublicKey != desc.SignerPublicKey {
285 + return false, false, errors.New("descriptor signer_public_key does not match pinned signer")
286 + }
287 +
288 + previousState := record.state.State
289 + previousDescriptor := record.state.Descriptor
290 + switch {
291 + case advertise:
292 + record.state.State = types.PeerStateAdvertised
293 + case record.state.State == types.PeerStateAdvertised:
294 + record.state.State = types.PeerStateAdvertised
295 + default:
296 + record.state.State = types.PeerStateVerified
297 + }
298 +
299 + record.state.Descriptor = desc
300 + record.state.LastSeenAt = now
301 + record.state.ConsecutiveFailures = 0
302 + r.peers[desc.RelayID] = record
303 +
304 + changed := added ||
305 + previousState != record.state.State ||
306 + !reflect.DeepEqual(previousDescriptor, record.state.Descriptor)
307 + return added, changed, nil
308 +}
309 +
310 +func (r *peerRegistry) registerDiscoveredPeers(targetRelayID, targetURL string, selfDescriptor types.RelayDescriptor, peerDescriptors []types.RelayDescriptor) (peerRegistrationResult, error) {
311 + if strings.TrimSpace(targetRelayID) == "" {
312 + return peerRegistrationResult{}, errors.New("target relay id is required")
313 + }
314 + if err := r.pin(targetRelayID, targetURL, selfDescriptor); err != nil {
315 + return peerRegistrationResult{}, err
316 + }
317 +
318 + added, changed, err := r.register(selfDescriptor, true)
319 + if err != nil {
320 + return peerRegistrationResult{}, err
321 + }
322 +
323 + result := peerRegistrationResult{
324 + updated: added || changed,
325 + peerSetChanged: changed,
326 + }
327 +
328 + for _, peerDescriptor := range peerDescriptors {
329 + hintAdded, hintChanged, err := r.register(peerDescriptor, false)
330 + if err != nil {
331 + return peerRegistrationResult{}, err
332 + }
333 + result.peerSetChanged = result.peerSetChanged || hintChanged
334 + if hintAdded || hintChanged {
335 + result.updated = true
336 + result.addedHintCount++
337 + }
338 + }
339 +
340 + return result, nil
341 +}
342 +
343 +func (r *peerRegistry) fail(relayID string) {
344 + if strings.TrimSpace(relayID) == "" {
345 + return
346 + }
347 +
348 + r.mu.Lock()
349 + defer r.mu.Unlock()
350 +
351 + record, ok := r.peers[relayID]
352 + if !ok {
353 + return
354 + }
355 + record.state.ConsecutiveFailures++
356 + r.peers[relayID] = record
357 +}
358 +
359 +func (r *peerRegistry) expire(relayID string) bool {
360 + if strings.TrimSpace(relayID) == "" {
361 + return false
362 + }
363 +
364 + r.mu.Lock()
365 + defer r.mu.Unlock()
366 +
367 + record, ok := r.peers[relayID]
368 + if !ok {
369 + return false
370 + }
371 + if record.state.State == types.PeerStateExpired {
372 + return false
373 + }
374 + record.state.State = types.PeerStateExpired
375 + r.peers[relayID] = record
376 + return true
377 +}
378 +
379 +func (r *peerRegistry) recordFailure(relayID string, err error, recoveryFailures int) peerFailureOutcome {
380 + r.fail(relayID)
381 +
382 + result := peerFailureOutcome{}
383 + state, _, ok := r.lookup(relayID)
384 + if ok &&
385 + state.State != types.PeerStateExpired &&
386 + state.ConsecutiveFailures >= recoveryFailures {
387 + if removed := r.expire(relayID); removed {
388 + result.expired = true
389 + result.expireReason = "recovery"
390 + result.consecutiveFailures = state.ConsecutiveFailures
391 + return result
392 + }
393 + }
394 +
395 + var apiErr *types.APIRequestError
396 + if errors.As(err, &apiErr) &&
397 + (apiErr.StatusCode == http.StatusForbidden ||
398 + apiErr.StatusCode == http.StatusNotFound ||
399 + apiErr.StatusCode == http.StatusGone) {
400 + if removed := r.expire(relayID); removed {
401 + result.expired = true
402 + result.expireReason = "status"
403 + }
404 + }
405 +
406 + return result
407 +}
portal/server.go
+176 -127
@@ -14,7 +14,6 @@ import (
14
15 "github.com/gosuda/keyless_tls/relay/l4"
16 "github.com/quic-go/quic-go"
17 - "github.com/rs/zerolog"
17 "github.com/rs/zerolog/log"
18 "golang.org/x/sync/errgroup"
19
@@ -84,7 +83,7 @@ type Server struct {
83 cfg ServerConfig
84 rootHost string
85 trustedProxyCIDRs []*net.IPNet
87 - discoveryCache *discovery.Cache
86 + peerRegistry *peerRegistry
87 shutdownOnce sync.Once
88 }
89
@@ -124,6 +123,9 @@ func NewServer(cfg ServerConfig) (*Server, error) {
123 if err != nil {
124 return nil, err
125 }
126 + if cfg.DiscoveryEnabled && strings.TrimSpace(wgConfig.PrivateKey) == "" {
127 + return nil, errors.New("wireguard private key is required when discovery is enabled")
128 + }
129
130 portMin, portMax := 0, 0
131 if cfg.UDPPortCount > 0 {
@@ -170,8 +172,8 @@ func NewServer(cfg ServerConfig) (*Server, error) {
172 }
173
174 if cfg.DiscoveryEnabled {
173 - s.discoveryCache = discovery.NewCache()
174 - if _, err := s.discoveryCache.UpsertSeedURLs(cfg.Bootstraps); err != nil {
175 + s.peerRegistry = newPeerRegistry()
176 + if _, err := s.peerRegistry.registerBootstrapURLs(cfg.Bootstraps); err != nil {
177 return nil, err
178 }
179 }
@@ -225,8 +227,8 @@ func (s *Server) Start(ctx context.Context, apiMux *http.ServeMux) error {
227
228 if s.wgConfig.PrivateKey != "" {
229 var snapshot map[string]types.PeerState
228 - if s.discoveryCache != nil {
229 - snapshot = s.discoveryCache.Snapshot()
230 + if s.peerRegistry != nil {
231 + snapshot = s.peerRegistry.snapshot()
232 }
233 peerMux := http.NewServeMux()
234 peerMux.HandleFunc(types.PathRoot, s.handleRoot)
@@ -318,8 +320,10 @@ func (s *Server) Shutdown(ctx context.Context) error {
320 shutdownErr = err
321 }
322 }
321 - if err := s.overlay.Shutdown(ctx); err != nil && shutdownErr == nil {
322 - shutdownErr = err
323 + if s.overlay != nil {
324 + if err := s.overlay.Shutdown(ctx); err != nil && shutdownErr == nil {
325 + shutdownErr = err
326 + }
327 }
328 if s.apiTLSClose != nil {
329 _ = s.apiTLSClose.Close()
@@ -555,141 +559,186 @@ func (s *Server) runQUICTunnelListener(listener *quic.Listener) error {
559 }
560 }
561
558 -func (s *Server) runDiscoveryLoop(ctx context.Context) error {
559 - ticker := time.NewTicker(defaultDiscoveryInterval)
560 - defer ticker.Stop()
562 +func (s *Server) applyDiscoveryResponse(targetRelayID, targetURL string, resp types.DiscoverResponse, requireSelfOverlay bool) (updated bool, addedHintCount int, warnErr error, err error) {
563 + if strings.TrimSpace(targetRelayID) == "" {
564 + return false, 0, nil, errors.New("target relay id is required")
565 + }
566
562 - for {
563 - peers := s.discoveryCache.KnownDescriptors()
564 - var overlayClient *http.Client
565 - if s.overlay != nil {
566 - overlayClient = s.overlay.Client()
567 + now := time.Now().UTC()
568 + selfDescriptor, peerDescriptors, warnErr := discovery.ValidateResponse(resp, now)
569 + if selfDescriptor.RelayID == "" {
570 + return false, 0, nil, errors.Join(warnErr, errors.New("discover response is missing self descriptor"))
571 + }
572 + if selfDescriptor.RelayID != targetRelayID {
573 + return false, 0, nil, errors.Join(warnErr, errors.New("discover response relay_id mismatch"))
574 + }
575 + if requireSelfOverlay {
576 + if err := requireOverlayPeerDescriptor(selfDescriptor); err != nil {
577 + return false, 0, nil, errors.Join(warnErr, err)
578 }
568 - for _, peer := range peers {
569 - discoverURL, discoverClient := peer.APIHTTPSAddr, (*http.Client)(nil)
570 - if state, pinned, _ := s.discoveryCache.Lookup(peer.RelayID); peer.SupportsOverlayPeer && overlayClient != nil && pinned && state.State != types.PeerStateExpired {
571 - if peer.OverlayIPv4 == "" {
572 - err := errors.New("relay peer is missing overlay ipv4")
573 - s.discoveryCache.RecordFailure(peer.RelayID)
574 - log.Warn().
575 - Err(err).
576 - Str("peer", peer.APIHTTPSAddr).
577 - Msg("discover peer failed")
578 - continue
579 - }
580 - discoverURL = "http://" + net.JoinHostPort(peer.OverlayIPv4, fmt.Sprintf("%d", wireguard.DefaultPeerAPIHTTPPort))
581 - discoverClient = overlayClient
582 - }
583 - resp, err := discovery.Discover(ctx, discoverURL, types.DiscoverRequest{}, nil, discoverClient)
584 - if err != nil {
585 - if ctx.Err() != nil {
586 - return nil
587 - }
588 - s.discoveryCache.RecordFailure(peer.RelayID)
589 - expireReason := ""
590 - consecutiveFailures := 0
591 - state, pinned, _ := s.discoveryCache.Lookup(peer.RelayID)
592 - if pinned &&
593 - state.State != types.PeerStateExpired &&
594 - peer.SupportsOverlayPeer &&
595 - state.ConsecutiveFailures >= defaultWGRecoveryFailures {
596 - if removed := s.discoveryCache.Expire(peer.RelayID); removed {
597 - expireReason = "recovery"
598 - consecutiveFailures = state.ConsecutiveFailures
599 - }
600 - }
601 - var apiErr *types.APIRequestError
602 - if expireReason == "" && errors.As(err, &apiErr) &&
603 - (apiErr.StatusCode == http.StatusForbidden ||
604 - apiErr.StatusCode == http.StatusNotFound ||
605 - apiErr.StatusCode == http.StatusGone) {
606 - if removed := s.discoveryCache.Expire(peer.RelayID); removed {
607 - expireReason = "status"
608 - }
609 - }
610 - if expireReason != "" {
611 - err = errors.Join(err, s.overlay.Sync(s.cfg.PortalURL, s.discoveryCache.Snapshot()))
612 - }
579 + }
580
614 - event := log.Warn().
615 - Err(err).
616 - Str("peer", peer.APIHTTPSAddr)
617 - if expireReason != "" {
618 - event = event.
619 - Bool("expired", true).
620 - Str("reason", expireReason)
621 - if consecutiveFailures > 0 {
622 - event = event.Int("consecutive_failures", consecutiveFailures)
623 - }
624 - }
625 - event.Msg("discover peer failed")
626 - continue
627 - }
581 + filteredPeerDescriptors := make([]types.RelayDescriptor, 0, len(peerDescriptors))
582
629 - now := time.Now().UTC()
630 - selfDescriptor, peerDescriptors, warnErr := discovery.ValidateResponse(resp, now)
631 - if selfDescriptor.RelayID == "" {
632 - err = errors.Join(warnErr, errors.New("discover response is missing self descriptor"))
583 + for _, peerDescriptor := range peerDescriptors {
584 + if err := requireOverlayPeerDescriptor(peerDescriptor); err != nil {
585 + warnErr = errors.Join(warnErr, fmt.Errorf("record hint %q: %w", peerDescriptor.RelayID, err))
586 + continue
587 + }
588 + filteredPeerDescriptors = append(filteredPeerDescriptors, peerDescriptor)
589 + }
590 +
591 + result, err := s.peerRegistry.registerDiscoveredPeers(targetRelayID, targetURL, selfDescriptor, filteredPeerDescriptors)
592 + if err != nil {
593 + return false, 0, nil, errors.Join(warnErr, err)
594 + }
595 + updated = result.updated
596 + addedHintCount = result.addedHintCount
597 +
598 + if result.peerSetChanged && s.overlay != nil {
599 + if err := s.overlay.Sync(s.cfg.PortalURL, s.peerRegistry.snapshot()); err != nil {
600 + warnErr = errors.Join(warnErr, err)
601 + }
602 + }
603 +
604 + return updated, addedHintCount, warnErr, nil
605 +}
606 +
607 +func (s *Server) runBootstrapDiscoveryPass(ctx context.Context) {
608 + for _, bootstrap := range s.peerRegistry.bootstrapPeers() {
609 + resp, err := discovery.Discover(ctx, bootstrap.APIHTTPSAddr, types.DiscoverRequest{}, nil, nil)
610 + if err != nil {
611 + if ctx.Err() != nil {
612 + return
613 }
634 - if err == nil {
635 - err = s.discoveryCache.PinIdentity(peer.RelayID, peer.APIHTTPSAddr, selfDescriptor)
614 + s.peerRegistry.fail(bootstrap.RelayID)
615 + log.Warn().
616 + Err(err).
617 + Str("peer", bootstrap.APIHTTPSAddr).
618 + Msg("bootstrap discovery failed")
619 + continue
620 + }
621 +
622 + updated, addedHintCount, warnErr, err := s.applyDiscoveryResponse(bootstrap.RelayID, bootstrap.APIHTTPSAddr, resp, false)
623 + if err != nil {
624 + s.peerRegistry.fail(bootstrap.RelayID)
625 + log.Warn().
626 + Err(err).
627 + Str("peer", bootstrap.APIHTTPSAddr).
628 + Msg("bootstrap discovery failed")
629 + continue
630 + }
631 +
632 + if updated || warnErr != nil {
633 + event := log.Info()
634 + if warnErr != nil {
635 + event = log.Warn().Err(warnErr)
636 }
637 - added, changed := false, false
638 - if err == nil {
639 - added, changed, err = s.discoveryCache.RecordVerified(selfDescriptor, true)
637 + event = event.
638 + Str("peer", bootstrap.APIHTTPSAddr).
639 + Int("bootstrap_count", len(s.peerRegistry.bootstrapPeers())).
640 + Int("peer_count", len(s.peerRegistry.syncablePeers())).
641 + Int("advertised_count", len(s.peerRegistry.advertisedPeers()))
642 + if addedHintCount > 0 {
643 + event = event.Int("added_hint_count", addedHintCount)
644 }
641 - if err != nil {
642 - s.discoveryCache.RecordFailure(peer.RelayID)
643 - log.Warn().
644 - Err(err).
645 - Str("peer", peer.APIHTTPSAddr).
646 - Msg("discover peer failed")
647 - continue
645 + if warnErr != nil {
646 + event.Msg("bootstrap discovery completed with warnings")
647 + } else {
648 + event.Msg("bootstrap discovery updated")
649 }
650 + }
651 + }
652 +}
653 +
654 +func (s *Server) runOverlayPeerDiscoveryPass(ctx context.Context) {
655 + if s.overlay == nil {
656 + return
657 + }
658 + overlayClient := s.overlay.Client()
659 + if overlayClient == nil {
660 + return
661 + }
662 +
663 + for _, peer := range s.peerRegistry.syncablePeers() {
664 + var failureErr error
665
650 - peerSetChanged := changed
651 - addedHintCount := 0
652 - for _, peerDescriptor := range peerDescriptors {
653 - hintAdded, hintChanged, err := s.discoveryCache.RecordVerified(peerDescriptor, false)
666 + if err := requireOverlayPeerDescriptor(peer); err != nil {
667 + failureErr = err
668 + } else {
669 + discoverURL := "http://" + net.JoinHostPort(peer.OverlayIPv4, fmt.Sprintf("%d", wireguard.DefaultPeerAPIHTTPPort))
670 + resp, err := discovery.Discover(ctx, discoverURL, types.DiscoverRequest{}, nil, overlayClient)
671 + if err != nil {
672 + if ctx.Err() != nil {
673 + return
674 + }
675 + failureErr = err
676 + } else {
677 + updated, addedHintCount, warnErr, err := s.applyDiscoveryResponse(peer.RelayID, peer.APIHTTPSAddr, resp, true)
678 if err != nil {
655 - warnErr = errors.Join(warnErr, fmt.Errorf("record hint %q: %w", peerDescriptor.RelayID, err))
679 + failureErr = err
680 + } else {
681 + if warnErr != nil {
682 + log.Warn().
683 + Err(warnErr).
684 + Str("peer", peer.APIHTTPSAddr).
685 + Int("bootstrap_count", len(s.peerRegistry.bootstrapPeers())).
686 + Int("peer_count", len(s.peerRegistry.syncablePeers())).
687 + Int("advertised_count", len(s.peerRegistry.advertisedPeers())).
688 + Int("added_hint_count", addedHintCount).
689 + Msg("overlay peer discovery completed with warnings")
690 + continue
691 + }
692 +
693 + if updated {
694 + event := log.Info().
695 + Str("peer", peer.APIHTTPSAddr).
696 + Int("bootstrap_count", len(s.peerRegistry.bootstrapPeers())).
697 + Int("peer_count", len(s.peerRegistry.syncablePeers())).
698 + Int("advertised_count", len(s.peerRegistry.advertisedPeers()))
699 + if addedHintCount > 0 {
700 + event = event.Int("added_hint_count", addedHintCount)
701 + }
702 + event.Msg("overlay peer discovery updated")
703 + }
704 continue
705 }
658 - peerSetChanged = peerSetChanged || hintChanged
659 - if hintAdded || hintChanged {
660 - addedHintCount++
661 - }
662 - }
663 - if peerSetChanged {
664 - if err := s.overlay.Sync(s.cfg.PortalURL, s.discoveryCache.Snapshot()); err != nil {
665 - warnErr = errors.Join(warnErr, err)
666 - }
706 }
707 + }
708
669 - updated := added || changed || addedHintCount > 0
670 - if updated || warnErr != nil {
671 - var event *zerolog.Event
672 - if warnErr != nil {
673 - event = log.Warn().
674 - Err(warnErr)
675 - } else {
676 - event = log.Info()
677 - }
678 - event.
679 - Str("peer", peer.APIHTTPSAddr).
680 - Bool("discoverable", selfDescriptor.SupportsOverlayPeer).
681 - Int("known_count", len(s.discoveryCache.KnownDescriptors())).
682 - Int("advertised_count", len(s.discoveryCache.AdvertisedDescriptors()))
683 - if addedHintCount > 0 {
684 - event.Int("added_hint_count", addedHintCount)
685 - }
686 - if updated {
687 - event.Msg("discovery peer updated")
688 - } else {
689 - event.Msg("discover peer completed with warnings")
690 - }
709 + result := s.peerRegistry.recordFailure(peer.RelayID, failureErr, defaultWGRecoveryFailures)
710 + if result.expired && s.overlay != nil {
711 + failureErr = errors.Join(failureErr, s.overlay.Sync(s.cfg.PortalURL, s.peerRegistry.snapshot()))
712 + }
713 +
714 + event := log.Warn().
715 + Err(failureErr).
716 + Str("peer", peer.APIHTTPSAddr)
717 + if result.expired {
718 + event = event.
719 + Bool("expired", true).
720 + Str("reason", result.expireReason)
721 + if result.consecutiveFailures > 0 {
722 + event = event.Int("consecutive_failures", result.consecutiveFailures)
723 }
724 }
725 + event.Msg("overlay peer discovery failed")
726 + }
727 +}
728 +
729 +func (s *Server) runDiscoveryLoop(ctx context.Context) error {
730 + ticker := time.NewTicker(defaultDiscoveryInterval)
731 + defer ticker.Stop()
732 +
733 + for {
734 + s.runBootstrapDiscoveryPass(ctx)
735 + if ctx.Err() != nil {
736 + return nil
737 + }
738 + s.runOverlayPeerDiscoveryPass(ctx)
739 + if ctx.Err() != nil {
740 + return nil
741 + }
742
743 select {
744 case <-ctx.Done():
portal/server_test.go
+85 -150
@@ -26,17 +26,33 @@ func mustSignedRelayDescriptor(t *testing.T, ownerPrivateKey, relayURL string) t
26 }
27
28 now := time.Now().UTC()
29 + wireGuardPrivateKey, err := utils.NormalizeWireGuardPrivateKey(strings.Repeat("44", 32))
30 + if err != nil {
31 + t.Fatalf("NormalizeWireGuardPrivateKey() error = %v", err)
32 + }
33 + wireGuardPublicKey, err := utils.WireGuardPublicKeyFromPrivate(wireGuardPrivateKey)
34 + if err != nil {
35 + t.Fatalf("WireGuardPublicKeyFromPrivate() error = %v", err)
36 + }
37 + overlayIPv4, err := utils.DeriveWireGuardOverlayIPv4(wireGuardPublicKey)
38 + if err != nil {
39 + t.Fatalf("DeriveWireGuardOverlayIPv4() error = %v", err)
40 + }
41 desc, err := discovery.SignedDescriptor(types.RelayDescriptor{
30 - RelayID: relayURL,
31 - OwnerAddress: identity.Address,
32 - SignerPublicKey: identity.PublicKey,
33 - Sequence: uint64(now.UnixMilli()),
34 - Version: 1,
35 - IssuedAt: now,
36 - ExpiresAt: now.Add(time.Hour),
37 - APIHTTPSAddr: relayURL,
38 - SupportsTCP: true,
39 - StatusState: "healthy",
42 + RelayID: relayURL,
43 + OwnerAddress: identity.Address,
44 + SignerPublicKey: identity.PublicKey,
45 + Sequence: uint64(now.UnixMilli()),
46 + Version: 1,
47 + IssuedAt: now,
48 + ExpiresAt: now.Add(time.Hour),
49 + APIHTTPSAddr: relayURL,
50 + WireGuardPublicKey: wireGuardPublicKey,
51 + WireGuardEndpoint: net.JoinHostPort(utils.PortalRootHost(relayURL), "51820"),
52 + OverlayIPv4: overlayIPv4,
53 + SupportsTCP: true,
54 + SupportsOverlayPeer: true,
55 + StatusState: "healthy",
56 }, identity.PrivateKey)
57 if err != nil {
58 t.Fatalf("SignedDescriptor() error = %v", err)
@@ -44,6 +60,21 @@ func mustSignedRelayDescriptor(t *testing.T, ownerPrivateKey, relayURL string) t
60 return desc
61 }
62
63 +func TestNewServerRequiresWireGuardWhenDiscoveryEnabled(t *testing.T) {
64 + t.Parallel()
65 +
66 + _, err := NewServer(ServerConfig{
67 + PortalURL: "https://portal.example.com",
68 + DiscoveryEnabled: true,
69 + })
70 + if err == nil {
71 + t.Fatal("NewServer() error = nil, want wireguard requirement error")
72 + }
73 + if !strings.Contains(err.Error(), "wireguard private key is required when discovery is enabled") {
74 + t.Fatalf("NewServer() error = %v, want wireguard requirement error", err)
75 + }
76 +}
77 +
78 func TestServerStartInitializesLocalACMEAndSigner(t *testing.T) {
79 t.Parallel()
80
@@ -266,135 +297,21 @@ func TestRegisterLeaseBuildsUDPEnabledRuntime(t *testing.T) {
297 }
298 }
299
269 -func TestServerStartServesOptionalDiscoveryRoutes(t *testing.T) {
270 - t.Parallel()
271 -
272 - ownerPrivateKey := strings.Repeat("11", 32)
273 - ownerIdentity, err := utils.ResolveSecp256k1Identity(ownerPrivateKey)
274 - if err != nil {
275 - t.Fatalf("ResolveSecp256k1Identity() error = %v", err)
276 - }
277 -
278 - server, err := NewServer(ServerConfig{
279 - PortalURL: "https://localhost:4017",
280 - OwnerPrivateKey: ownerPrivateKey,
281 - Bootstraps: []string{"https://bootstrap.example.com"},
282 - ACME: acme.Config{KeyDir: t.TempDir()},
283 - APIListenAddr: "127.0.0.1:0",
284 - SNIListenAddr: "127.0.0.1:0",
285 - DiscoveryEnabled: true,
286 - })
287 - if err != nil {
288 - t.Fatalf("NewServer() error = %v", err)
289 - }
290 - registerResp, err := server.registerLease(types.RegisterRequest{
291 - Name: "demo",
292 - ReverseToken: "tok_demo",
293 - Bootstraps: []string{"https://relay-a.example.com", "https://bootstrap.example.com"},
294 - }, "203.0.113.10")
295 - if err != nil {
296 - t.Fatalf("registerLease() error = %v", err)
297 - }
298 - if !reflect.DeepEqual(registerResp.Bootstraps, []string{server.PortalURL()}) {
299 - t.Fatalf("registerLease() bootstraps = %v, want [%q]", registerResp.Bootstraps, server.PortalURL())
300 - }
301 -
302 - ctx, cancel := context.WithCancel(context.Background())
303 - defer cancel()
304 -
305 - if err := server.Start(ctx, nil); err != nil {
306 - t.Fatalf("Start() error = %v", err)
307 - }
308 -
309 - client := &http.Client{
310 - Transport: &http.Transport{
311 - TLSClientConfig: &tls.Config{InsecureSkipVerify: true},
312 - },
313 - }
314 - t.Cleanup(func() {
315 - client.CloseIdleConnections()
316 - cancel()
317 - if err := server.Wait(); err != nil {
318 - t.Fatalf("Wait() error = %v", err)
319 - }
320 - })
321 -
322 - resp, err := client.Get("https://" + utils.HostPortOrLoopback(server.APIAddr()) + types.PathDiscovery + "?root_host=localhost&name=demo")
323 - if err != nil {
324 - t.Fatalf("GET discovery resolve error = %v", err)
325 - }
326 - defer resp.Body.Close()
327 -
328 - if resp.StatusCode != http.StatusOK {
329 - t.Fatalf("GET discovery resolve status = %d, want %d", resp.StatusCode, http.StatusOK)
330 - }
331 -
332 - var envelope types.APIEnvelope[types.DiscoverResponse]
333 - if err := json.NewDecoder(resp.Body).Decode(&envelope); err != nil {
334 - t.Fatalf("decode discovery resolve response: %v", err)
335 - }
336 - if !envelope.OK {
337 - t.Fatalf("discovery resolve envelope = %+v, want ok", envelope)
338 - }
339 - if envelope.Data.ProtocolVersion != 1 {
340 - t.Fatalf("resolve protocol_version = %d, want 1", envelope.Data.ProtocolVersion)
341 - }
342 - if envelope.Data.GeneratedAt.IsZero() {
343 - t.Fatal("resolve generated_at = zero, want timestamp")
344 - }
345 - if envelope.Data.Service == nil || !envelope.Data.Service.Found {
346 - t.Fatalf("resolve service = %+v, want found=true", envelope.Data.Service)
347 - }
348 - if envelope.Data.Service.OwnerAddress != ownerIdentity.Address {
349 - t.Fatalf("resolve service owner address = %q, want relay owner address", envelope.Data.Service.OwnerAddress)
350 - }
351 - if envelope.Data.Service.Hostname != "demo.localhost" {
352 - t.Fatalf("resolve service hostname = %q, want %q", envelope.Data.Service.Hostname, "demo.localhost")
353 - }
354 - if envelope.Data.Service.RelayID != envelope.Data.Self.RelayID {
355 - t.Fatalf("resolve service relay_id = %q, want %q", envelope.Data.Service.RelayID, envelope.Data.Self.RelayID)
356 - }
357 -
358 - if _, err := discovery.ValidateDescriptor(envelope.Data.Self, time.Now().UTC()); err != nil {
359 - t.Fatalf("ValidateDescriptor(self) error = %v", err)
360 - }
361 - relayURLs := make([]string, 0, 1+len(envelope.Data.Peers))
362 - relayURLs = append(relayURLs, envelope.Data.Self.APIHTTPSAddr)
363 - for _, peer := range envelope.Data.Peers {
364 - if strings.TrimSpace(peer.APIHTTPSAddr) != "" {
365 - relayURLs = append(relayURLs, peer.APIHTTPSAddr)
366 - }
367 - }
368 - if !reflect.DeepEqual(relayURLs, []string{server.PortalURL()}) {
369 - t.Fatalf("resolve relay urls = %v, want [%q]", relayURLs, server.PortalURL())
370 - }
371 - if envelope.Data.Self.OwnerAddress != ownerIdentity.Address {
372 - t.Fatalf("self relay owner address = %q, want %q", envelope.Data.Self.OwnerAddress, ownerIdentity.Address)
373 - }
374 - if envelope.Data.Self.SignerPublicKey != ownerIdentity.PublicKey {
375 - t.Fatalf("self relay signer public key = %q, want %q", envelope.Data.Self.SignerPublicKey, ownerIdentity.PublicKey)
376 - }
377 - if envelope.Data.Self.StatusState != "healthy" {
378 - t.Fatalf("self relay status_state = %q, want %q", envelope.Data.Self.StatusState, "healthy")
379 - }
380 - if !server.DiscoveryEnabled() {
381 - t.Fatal("DiscoveryEnabled() = false, want true")
382 - }
383 -}
384 -
300 func TestServerUpsertDiscoverySeedURLsSkipsLocalRelayHosts(t *testing.T) {
301 t.Parallel()
302
303 server, err := NewServer(ServerConfig{
389 - PortalURL: "https://portal.example.com",
390 - Bootstraps: []string{"https://bootstrap.example.com"},
391 - DiscoveryEnabled: true,
304 + PortalURL: "https://portal.example.com",
305 + Bootstraps: []string{"https://bootstrap.example.com"},
306 + WireGuardPrivateKey: strings.Repeat("23", 32),
307 + DiscoveryPort: 41022,
308 + DiscoveryEnabled: true,
309 })
310 if err != nil {
311 t.Fatalf("NewServer() error = %v", err)
312 }
313
397 - added, err := server.discoveryCache.UpsertSeedURLs([]string{
314 + added, err := server.peerRegistry.registerBootstrapURLs([]string{
315 "https://localhost:4017",
316 "https://relay-a.example.com",
317 "https://127.0.0.1:4017",
@@ -410,8 +327,8 @@ func TestServerUpsertDiscoverySeedURLsSkipsLocalRelayHosts(t *testing.T) {
327 if err != nil {
328 t.Fatalf("ExcludeLocalRelayURLs() error = %v", err)
329 }
413 - knownURLs := make([]string, 0, len(server.discoveryCache.KnownDescriptors()))
414 - for _, descriptor := range server.discoveryCache.KnownDescriptors() {
330 + knownURLs := make([]string, 0, len(server.peerRegistry.bootstrapPeers()))
331 + for _, descriptor := range server.peerRegistry.bootstrapPeers() {
332 if apiURL := strings.TrimSpace(descriptor.APIHTTPSAddr); apiURL != "" {
333 knownURLs = append(knownURLs, apiURL)
334 }
@@ -421,10 +338,13 @@ func TestServerUpsertDiscoverySeedURLsSkipsLocalRelayHosts(t *testing.T) {
338 t.Fatalf("ExcludeLocalRelayURLs() known error = %v", err)
339 }
340 if !reflect.DeepEqual(knownURLs, knownRelayURLs) {
424 - t.Fatalf("KnownDescriptors() = %v, want [%q %q]", knownURLs, "https://bootstrap.example.com", "https://relay-a.example.com")
341 + t.Fatalf("BootstrapDescriptors() = %v, want [%q %q]", knownURLs, "https://bootstrap.example.com", "https://relay-a.example.com")
342 + }
343 + if len(server.peerRegistry.syncablePeers()) != 0 {
344 + t.Fatalf("syncablePeers() = %v, want empty before direct confirmation", server.peerRegistry.syncablePeers())
345 }
426 - if len(server.discoveryCache.AdvertisedDescriptors()) != 0 {
427 - t.Fatalf("AdvertisedDescriptors() = %v, want empty before direct confirmation", server.discoveryCache.AdvertisedDescriptors())
346 + if len(server.peerRegistry.advertisedPeers()) != 0 {
347 + t.Fatalf("advertisedPeers() = %v, want empty before direct confirmation", server.peerRegistry.advertisedPeers())
348 }
349 }
350
@@ -433,10 +353,12 @@ func TestServerRecordVerifiedDiscoveryPeerRequiresDirectConfirmation(t *testing.
353
354 ownerPrivateKey := strings.Repeat("11", 32)
355 server, err := NewServer(ServerConfig{
436 - PortalURL: "https://portal.example.com",
437 - Bootstraps: []string{"https://bootstrap.example.com"},
438 - OwnerPrivateKey: ownerPrivateKey,
439 - DiscoveryEnabled: true,
356 + PortalURL: "https://portal.example.com",
357 + Bootstraps: []string{"https://bootstrap.example.com"},
358 + OwnerPrivateKey: ownerPrivateKey,
359 + WireGuardPrivateKey: strings.Repeat("24", 32),
360 + DiscoveryPort: 41023,
361 + DiscoveryEnabled: true,
362 })
363 if err != nil {
364 t.Fatalf("NewServer() error = %v", err)
@@ -445,7 +367,7 @@ func TestServerRecordVerifiedDiscoveryPeerRequiresDirectConfirmation(t *testing.
367 bootstrapDesc := mustSignedRelayDescriptor(t, ownerPrivateKey, "https://bootstrap.example.com")
368 relayADesc := mustSignedRelayDescriptor(t, ownerPrivateKey, "https://relay-a.example.com")
369
448 - added, changed, err := server.discoveryCache.RecordVerified(bootstrapDesc, true)
370 + added, changed, err := server.peerRegistry.register(bootstrapDesc, true)
371 if err != nil {
372 t.Fatalf("RecordVerified() error = %v", err)
373 }
@@ -456,7 +378,7 @@ func TestServerRecordVerifiedDiscoveryPeerRequiresDirectConfirmation(t *testing.
378 t.Fatal("RecordVerified() changed = false, want true")
379 }
380
459 - added, changed, err = server.discoveryCache.RecordVerified(relayADesc, false)
381 + added, changed, err = server.peerRegistry.register(relayADesc, false)
382 if err != nil {
383 t.Fatalf("RecordVerified() hinted error = %v", err)
384 }
@@ -466,8 +388,8 @@ func TestServerRecordVerifiedDiscoveryPeerRequiresDirectConfirmation(t *testing.
388 if !changed {
389 t.Fatal("RecordVerified() hinted changed = false, want true")
390 }
469 - knownURLs := make([]string, 0, len(server.discoveryCache.KnownDescriptors()))
470 - for _, descriptor := range server.discoveryCache.KnownDescriptors() {
391 + knownURLs := make([]string, 0, len(server.peerRegistry.bootstrapPeers()))
392 + for _, descriptor := range server.peerRegistry.bootstrapPeers() {
393 if apiURL := strings.TrimSpace(descriptor.APIHTTPSAddr); apiURL != "" {
394 knownURLs = append(knownURLs, apiURL)
395 }
@@ -476,11 +398,24 @@ func TestServerRecordVerifiedDiscoveryPeerRequiresDirectConfirmation(t *testing.
398 if err != nil {
399 t.Fatalf("ExcludeLocalRelayURLs() known error = %v", err)
400 }
479 - if !reflect.DeepEqual(knownURLs, []string{"https://bootstrap.example.com", "https://relay-a.example.com"}) {
480 - t.Fatalf("KnownDescriptors() = %v, want [%q %q]", knownURLs, "https://bootstrap.example.com", "https://relay-a.example.com")
401 + if !reflect.DeepEqual(knownURLs, []string{"https://bootstrap.example.com"}) {
402 + t.Fatalf("BootstrapDescriptors() = %v, want [%q]", knownURLs, "https://bootstrap.example.com")
403 + }
404 + syncableURLs := make([]string, 0, len(server.peerRegistry.syncablePeers()))
405 + for _, descriptor := range server.peerRegistry.syncablePeers() {
406 + if apiURL := strings.TrimSpace(descriptor.APIHTTPSAddr); apiURL != "" {
407 + syncableURLs = append(syncableURLs, apiURL)
408 + }
409 + }
410 + syncableURLs, err = utils.ExcludeLocalRelayURLs(syncableURLs...)
411 + if err != nil {
412 + t.Fatalf("ExcludeLocalRelayURLs() syncable error = %v", err)
413 + }
414 + if !reflect.DeepEqual(syncableURLs, []string{"https://relay-a.example.com"}) {
415 + t.Fatalf("SyncablePeerDescriptors() = %v, want [%q]", syncableURLs, "https://relay-a.example.com")
416 }
482 - advertisedURLs := make([]string, 0, len(server.discoveryCache.AdvertisedDescriptors()))
483 - for _, descriptor := range server.discoveryCache.AdvertisedDescriptors() {
417 + advertisedURLs := make([]string, 0, len(server.peerRegistry.advertisedPeers()))
418 + for _, descriptor := range server.peerRegistry.advertisedPeers() {
419 if apiURL := strings.TrimSpace(descriptor.APIHTTPSAddr); apiURL != "" {
420 advertisedURLs = append(advertisedURLs, apiURL)
421 }
@@ -493,7 +428,7 @@ func TestServerRecordVerifiedDiscoveryPeerRequiresDirectConfirmation(t *testing.
428 t.Fatalf("AdvertisedDescriptors() = %v, want [%q]", advertisedURLs, "https://bootstrap.example.com")
429 }
430
496 - snapshot := server.discoveryCache.Snapshot()
431 + snapshot := server.peerRegistry.snapshot()
432 if snapshot[bootstrapDesc.RelayID].State != types.PeerStateAdvertised {
433 t.Fatalf("bootstrap state = %q, want %q", snapshot[bootstrapDesc.RelayID].State, types.PeerStateAdvertised)
434 }
@@ -501,7 +436,7 @@ func TestServerRecordVerifiedDiscoveryPeerRequiresDirectConfirmation(t *testing.
436 t.Fatalf("relay-a state = %q, want %q", snapshot[relayADesc.RelayID].State, types.PeerStateVerified)
437 }
438
504 - added, changed, err = server.discoveryCache.RecordVerified(relayADesc, true)
439 + added, changed, err = server.peerRegistry.register(relayADesc, true)
440 if err != nil {
441 t.Fatalf("RecordVerified() second error = %v", err)
442 }
@@ -512,7 +447,7 @@ func TestServerRecordVerifiedDiscoveryPeerRequiresDirectConfirmation(t *testing.
447 t.Fatal("RecordVerified() second changed = false, want true")
448 }
449 advertisedURLs = advertisedURLs[:0]
515 - for _, descriptor := range server.discoveryCache.AdvertisedDescriptors() {
450 + for _, descriptor := range server.peerRegistry.advertisedPeers() {
451 if apiURL := strings.TrimSpace(descriptor.APIHTTPSAddr); apiURL != "" {
452 advertisedURLs = append(advertisedURLs, apiURL)
453 }
sdk/api_client.go
+25 -7
@@ -94,7 +94,11 @@ func (a *apiClient) close() {
94 }
95 }
96
97 -func (a *apiClient) registerLease(ctx context.Context, ttl time.Duration, udpEnabled bool, bootstraps []string) (types.RegisterResponse, error) {
97 +func (a *apiClient) registerLease(ctx context.Context, ttl time.Duration, udpEnabled bool) (types.RegisterResponse, error) {
98 + if err := a.ensureHTTPClient(ctx); err != nil {
99 + return types.RegisterResponse{}, err
100 + }
101 +
102 var resp types.RegisterResponse
103 if err := a.doJSON(ctx, http.MethodPost, types.PathSDKRegister, types.RegisterRequest{
104 Name: a.name,
@@ -102,16 +106,15 @@ func (a *apiClient) registerLease(ctx context.Context, ttl time.Duration, udpEna
106 OwnerAddress: a.ownerAddress,
107 ReverseToken: a.reverseToken,
108 TTL: int(ttl / time.Second),
105 - Bootstraps: bootstraps,
109 UDPEnabled: udpEnabled,
107 - ReportedIP: a.resolvedPublicIP,
110 + ReportedIP: a.reportedIP(ctx),
111 }, &resp); err != nil {
112 return types.RegisterResponse{}, err
113 }
114 return resp, nil
115 }
116
114 -func (a *apiClient) ensureReady(ctx context.Context) error {
117 +func (a *apiClient) ensureHTTPClient(ctx context.Context) error {
118 if a.httpClient != nil && a.rawTLSConfig != nil {
119 return nil
120 }
@@ -146,11 +149,14 @@ func (a *apiClient) ensureReady(ctx context.Context) error {
149 a.httpClient = httpClient
150 a.rawTLSConfig = rawTLSConfig
151
152 + return nil
153 +}
154 +
155 +func (a *apiClient) reportedIP(ctx context.Context) string {
156 if a.resolvedPublicIP == "" {
157 a.resolvedPublicIP = utils.ResolvePublicIP(ctx)
158 }
152 -
153 - return nil
159 + return a.resolvedPublicIP
160 }
161
162 func (a *apiClient) ensureCompatible(ctx context.Context, httpClient *http.Client) error {
@@ -174,11 +180,15 @@ func (a *apiClient) ensureCompatible(ctx context.Context, httpClient *http.Clien
180 }
181
182 func (a *apiClient) renewLease(ctx context.Context, leaseID string, ttl time.Duration) error {
183 + if err := a.ensureHTTPClient(ctx); err != nil {
184 + return err
185 + }
186 +
187 return a.doJSON(ctx, http.MethodPost, types.PathSDKRenew, types.RenewRequest{
188 LeaseID: leaseID,
189 ReverseToken: a.reverseToken,
190 TTL: int(ttl / time.Second),
181 - ReportedIP: a.resolvedPublicIP,
191 + ReportedIP: a.reportedIP(ctx),
192 }, &types.RenewResponse{})
193 }
194
@@ -190,6 +200,10 @@ func (a *apiClient) unregisterLease(ctx context.Context, leaseID string) error {
200 }
201
202 func (a *apiClient) openReverseSession(ctx context.Context, leaseID string) (net.Conn, error) {
203 + if err := a.ensureHTTPClient(ctx); err != nil {
204 + return nil, err
205 + }
206 +
207 dialer := &tls.Dialer{
208 NetDialer: &net.Dialer{Timeout: a.dialTimeout},
209 Config: a.rawTLSConfig.Clone(),
@@ -310,6 +324,10 @@ func (c *bufferedConn) Read(p []byte) (int, error) {
324
325 // openQUICSession opens a QUIC connection to the relay for datagram transport.
326 func (a *apiClient) openQUICSession(ctx context.Context, leaseID, reverseToken string) (*quic.Conn, error) {
327 + if err := a.ensureHTTPClient(ctx); err != nil {
328 + return nil, err
329 + }
330 +
331 tlsConf := a.rawTLSConfig.Clone()
332 tlsConf.NextProtos = []string{"portal-tunnel"}
333
sdk/expose.go
+278 -164
@@ -38,16 +38,22 @@ type Exposure struct {
38 accepted chan net.Conn
39 datagrams chan types.DatagramFrame
40
41 - mu sync.RWMutex
42 - knownRelayURLs []string
43 - activeRelayURLs []string
44 - bannedRelayURLs []string
45 - listeners map[string]*Listener
41 + mu sync.RWMutex
42 + knownRelayURLs []string
43 + bannedRelayURLs []string
44 + discoveryPins map[string]discoveryIdentity
45 + discoveryRelayIDsByURL map[string]string
46 + listeners map[string]*Listener
47
48 closeOnce sync.Once
49 connSeq atomic.Uint64
50 }
51
52 +type discoveryIdentity struct {
53 + apiURL string
54 + signerPublicKey string
55 +}
56 +
57 type ExposeConfig struct {
58 RelayURLs []string
59 Name string
@@ -88,21 +94,23 @@ func Expose(ctx context.Context, cfg ExposeConfig) (*Exposure, error) {
94
95 exposureCtx, cancel := context.WithCancel(ctx)
96 exposure := &Exposure{
91 - cancel: cancel,
92 - done: exposureCtx.Done(),
93 - name: cfg.Name,
94 - TargetAddr: targetAddr,
95 - UDPAddr: udpAddr,
96 - reverseToken: cfg.ReverseToken,
97 - udpEnabled: cfg.UDPEnabled,
98 - banMITM: cfg.BanMITM,
99 - metadata: cfg.Metadata.Copy(),
100 - ownerAddress: identity.Address,
101 - rootCAPEM: append([]byte(nil), cfg.RootCAPEM...),
102 - discoveryEnabled: cfg.Discovery,
103 - accepted: make(chan net.Conn, max(len(relayURLs)*defaultReadyTarget*2, 1)),
104 - datagrams: make(chan types.DatagramFrame, max(len(relayURLs)*32, 1)),
105 - listeners: make(map[string]*Listener, len(relayURLs)),
97 + cancel: cancel,
98 + done: exposureCtx.Done(),
99 + name: cfg.Name,
100 + TargetAddr: targetAddr,
101 + UDPAddr: udpAddr,
102 + reverseToken: cfg.ReverseToken,
103 + udpEnabled: cfg.UDPEnabled,
104 + banMITM: cfg.BanMITM,
105 + metadata: cfg.Metadata.Copy(),
106 + ownerAddress: identity.Address,
107 + rootCAPEM: append([]byte(nil), cfg.RootCAPEM...),
108 + discoveryEnabled: cfg.Discovery,
109 + accepted: make(chan net.Conn, max(len(relayURLs)*defaultReadyTarget*2, 1)),
110 + datagrams: make(chan types.DatagramFrame, max(len(relayURLs)*32, 1)),
111 + discoveryPins: make(map[string]discoveryIdentity),
112 + discoveryRelayIDsByURL: make(map[string]string),
113 + listeners: make(map[string]*Listener, len(relayURLs)),
114 }
115
116 if len(relayURLs) > 0 {
@@ -125,7 +133,7 @@ func Expose(ctx context.Context, cfg ExposeConfig) (*Exposure, error) {
133
134 if len(relayURLs) > 0 {
135 exposure.mu.RLock()
128 - activeRelayURLs := append([]string(nil), exposure.activeRelayURLs...)
136 + activeRelayURLs, _ := exposure.activeSnapshotLocked()
137 exposure.mu.RUnlock()
138 log.Info().
139 Str("release_version", types.ReleaseVersion).
@@ -137,37 +145,45 @@ func Expose(ctx context.Context, cfg ExposeConfig) (*Exposure, error) {
145 return exposure, nil
146 }
147
140 -func (e *Exposure) KnownRelayURLs() []string {
148 +func (e *Exposure) ActiveRelayURLs() []string {
149 e.mu.RLock()
150 defer e.mu.RUnlock()
151
144 - if len(e.knownRelayURLs) == 0 {
152 + if len(e.knownRelayURLs) == 0 || len(e.listeners) == 0 {
153 return nil
154 }
155
148 - return append([]string(nil), e.knownRelayURLs...)
149 -}
150 -
151 -func (e *Exposure) ActiveRelayURLs() []string {
152 - e.mu.RLock()
153 - defer e.mu.RUnlock()
154 -
155 - if len(e.activeRelayURLs) == 0 {
156 + activeRelayURLs := make([]string, 0, len(e.listeners))
157 + for _, relayURL := range e.knownRelayURLs {
158 + if _, ok := e.listeners[relayURL]; ok {
159 + activeRelayURLs = append(activeRelayURLs, relayURL)
160 + }
161 + }
162 + if len(activeRelayURLs) == 0 {
163 return nil
164 }
158 -
159 - return append([]string(nil), e.activeRelayURLs...)
165 + return activeRelayURLs
166 }
167
162 -func (e *Exposure) BannedRelayURLs() []string {
163 - e.mu.RLock()
164 - defer e.mu.RUnlock()
165 -
166 - if len(e.bannedRelayURLs) == 0 {
167 - return nil
168 +func (e *Exposure) activeSnapshotLocked() ([]string, []*Listener) {
169 + if len(e.knownRelayURLs) == 0 || len(e.listeners) == 0 {
170 + return nil, nil
171 }
172
170 - return append([]string(nil), e.bannedRelayURLs...)
173 + activeRelayURLs := make([]string, 0, len(e.listeners))
174 + listeners := make([]*Listener, 0, len(e.listeners))
175 + for _, relayURL := range e.knownRelayURLs {
176 + listener, ok := e.listeners[relayURL]
177 + if !ok {
178 + continue
179 + }
180 + activeRelayURLs = append(activeRelayURLs, relayURL)
181 + listeners = append(listeners, listener)
182 + }
183 + if len(activeRelayURLs) == 0 {
184 + return nil, nil
185 + }
186 + return activeRelayURLs, listeners
187 }
188
189 func (e *Exposure) Accept() (net.Conn, error) {
@@ -202,9 +218,9 @@ func (e *Exposure) Addr() net.Addr {
218 func (e *Exposure) RunHTTP(ctx context.Context, handler http.Handler, localAddr string) error {
219 var relayListener net.Listener
220 e.mu.RLock()
205 - hasActiveRelays := len(e.activeRelayURLs) > 0
221 + _, activeListeners := e.activeSnapshotLocked()
222 e.mu.RUnlock()
207 - if hasActiveRelays {
223 + if len(activeListeners) > 0 {
224 relayListener = e
225 }
226 return RunHTTP(ctx, relayListener, handler, localAddr)
@@ -319,12 +335,11 @@ func (e *Exposure) Close() error {
335 }
336
337 e.mu.RLock()
338 + relayURLs := make([]string, 0, len(e.listeners))
339 listeners := make([]*Listener, 0, len(e.listeners))
323 - activeRelayURLs := append([]string(nil), e.activeRelayURLs...)
324 - for _, relayURL := range e.activeRelayURLs {
325 - if listener, ok := e.listeners[relayURL]; ok {
326 - listeners = append(listeners, listener)
327 - }
340 + for relayURL, listener := range e.listeners {
341 + relayURLs = append(relayURLs, relayURL)
342 + listeners = append(listeners, listener)
343 }
344 e.mu.RUnlock()
345
@@ -336,12 +351,12 @@ func (e *Exposure) Close() error {
351
352 event := log.Info().
353 Int("relay_count", len(listeners)).
339 - Strs("relays", activeRelayURLs)
354 + Strs("relays", relayURLs)
355 if closeErr != nil {
356 event = log.Warn().
357 Err(closeErr).
358 Int("relay_count", len(listeners)).
344 - Strs("relays", activeRelayURLs)
359 + Strs("relays", relayURLs)
360 }
361 event.Msg("exposure closed")
362 })
@@ -353,12 +368,20 @@ func (e *Exposure) setRelayURLs(relayURLs []string, failOnError bool) ([]string,
368 return nil, nil
369 }
370
356 - relayURLs = utils.FilterRelayURLs(append([]string(nil), relayURLs...), e.bannedRelayURLs)
371 + e.mu.RLock()
372 + bannedRelayURLs := append([]string(nil), e.bannedRelayURLs...)
373 + e.mu.RUnlock()
374 + relayURLs = utils.FilterRelayURLs(append([]string(nil), relayURLs...), bannedRelayURLs)
375 +
376 e.mu.Lock()
377 existing := make(map[string]struct{}, len(e.knownRelayURLs))
378 for _, relayURL := range e.knownRelayURLs {
379 existing[relayURL] = struct{}{}
380 }
381 + desired := make(map[string]struct{}, len(relayURLs))
382 + for _, relayURL := range relayURLs {
383 + desired[relayURL] = struct{}{}
384 + }
385
386 added := make([]string, 0, len(relayURLs))
387 missing := make([]string, 0)
@@ -370,10 +393,31 @@ func (e *Exposure) setRelayURLs(relayURLs []string, failOnError bool) ([]string,
393 missing = append(missing, relayURL)
394 }
395 }
396 + removed := make([]string, 0)
397 + staleListeners := make([]*Listener, 0)
398 + for relayURL, listener := range e.listeners {
399 + if _, ok := desired[relayURL]; ok {
400 + continue
401 + }
402 + removed = append(removed, relayURL)
403 + staleListeners = append(staleListeners, listener)
404 + delete(e.listeners, relayURL)
405 + }
406 e.knownRelayURLs = append([]string(nil), relayURLs...)
374 - e.activeRelayURLs = append([]string(nil), relayURLs...)
407 e.mu.Unlock()
408
409 + for i, listener := range staleListeners {
410 + if listener == nil {
411 + continue
412 + }
413 + if err := listener.Close(); err != nil && !errors.Is(err, net.ErrClosed) {
414 + log.Warn().Err(err).Str("relay_url", removed[i]).Msg("close stale relay listener")
415 + }
416 + }
417 + if len(removed) > 0 {
418 + log.Info().Int("removed_count", len(removed)).Strs("removed_relays", removed).Msg("stale relays removed from exposure")
419 + }
420 +
421 for _, relayURL := range missing {
422 listener, err := e.newListener(relayURL)
423 if err != nil {
@@ -388,10 +432,63 @@ func (e *Exposure) setRelayURLs(relayURLs []string, failOnError bool) ([]string,
432 return added, nil
433 }
434
435 +func (e *Exposure) pinDiscoverySelfDescriptor(targetURL string, desc types.RelayDescriptor) error {
436 + normalizedTargetURL, err := utils.NormalizeRelayURL(targetURL)
437 + if err != nil {
438 + return err
439 + }
440 + if desc.APIHTTPSAddr != normalizedTargetURL {
441 + return errors.New("descriptor api_https_addr does not match target url")
442 + }
443 + return e.pinDiscoveredDescriptor(desc)
444 +}
445 +
446 +func (e *Exposure) pinDiscoveredDescriptor(desc types.RelayDescriptor) error {
447 + relayID := strings.TrimSpace(desc.RelayID)
448 + apiURL := strings.TrimSpace(desc.APIHTTPSAddr)
449 + signerPublicKey := strings.TrimSpace(desc.SignerPublicKey)
450 + if relayID == "" {
451 + return errors.New("descriptor relay_id is required")
452 + }
453 + if apiURL == "" {
454 + return errors.New("descriptor api_https_addr is required")
455 + }
456 + if signerPublicKey == "" {
457 + return errors.New("descriptor signer_public_key is required")
458 + }
459 +
460 + e.mu.Lock()
461 + defer e.mu.Unlock()
462 + if e.discoveryPins == nil {
463 + e.discoveryPins = make(map[string]discoveryIdentity)
464 + }
465 + if e.discoveryRelayIDsByURL == nil {
466 + e.discoveryRelayIDsByURL = make(map[string]string)
467 + }
468 +
469 + if pinned, ok := e.discoveryPins[relayID]; ok {
470 + if pinned.apiURL != apiURL {
471 + return errors.New("descriptor api_https_addr does not match pinned relay url")
472 + }
473 + if pinned.signerPublicKey != signerPublicKey {
474 + return errors.New("descriptor signer_public_key does not match pinned signer")
475 + }
476 + }
477 + if pinnedRelayID, ok := e.discoveryRelayIDsByURL[apiURL]; ok && pinnedRelayID != relayID {
478 + return errors.New("descriptor relay_id does not match pinned relay url identity")
479 + }
480 +
481 + e.discoveryPins[relayID] = discoveryIdentity{
482 + apiURL: apiURL,
483 + signerPublicKey: signerPublicKey,
484 + }
485 + e.discoveryRelayIDsByURL[apiURL] = relayID
486 + return nil
487 +}
488 +
489 func (e *Exposure) banRelayURL(relayURL string) {
490 e.mu.Lock()
491 e.knownRelayURLs = utils.RemoveRelayURL(e.knownRelayURLs, relayURL)
394 - e.activeRelayURLs = utils.RemoveRelayURL(e.activeRelayURLs, relayURL)
492 e.bannedRelayURLs = utils.AppendUniqueRelayURL(e.bannedRelayURLs, relayURL)
493 delete(e.listeners, relayURL)
494 bannedRelayURLs := append([]string(nil), e.bannedRelayURLs...)
@@ -404,24 +501,14 @@ func (e *Exposure) banRelayURL(relayURL string) {
501 }
502
503 func (e *Exposure) newListener(relayURL string) (*Listener, error) {
407 - bootstraps := []string(nil)
408 - if e.discoveryEnabled {
409 - e.mu.RLock()
410 - if len(e.knownRelayURLs) > 0 {
411 - bootstraps = append([]string(nil), e.knownRelayURLs...)
412 - }
413 - e.mu.RUnlock()
414 - }
415 -
504 cfg := ListenerConfig{
417 - Name: e.name,
418 - ReverseToken: e.reverseToken,
419 - UDPEnabled: e.udpEnabled,
420 - BanMITM: e.banMITM,
421 - RegisterBootstraps: bootstraps,
422 - Metadata: e.metadata.Copy(),
423 - RootCAPEM: append([]byte(nil), e.rootCAPEM...),
424 - ownerAddress: e.ownerAddress,
505 + Name: e.name,
506 + ReverseToken: e.reverseToken,
507 + UDPEnabled: e.udpEnabled,
508 + BanMITM: e.banMITM,
509 + Metadata: e.metadata.Copy(),
510 + RootCAPEM: append([]byte(nil), e.rootCAPEM...),
511 + ownerAddress: e.ownerAddress,
512 }
513 return NewListener(context.Background(), relayURL, cfg)
514 }
@@ -591,12 +678,7 @@ func (e *Exposure) WaitDatagramReady(ctx context.Context) ([]string, error) {
678
679 for {
680 e.mu.RLock()
594 - listeners := make([]*Listener, 0, len(e.listeners))
595 - for _, relayURL := range e.activeRelayURLs {
596 - if listener, ok := e.listeners[relayURL]; ok {
597 - listeners = append(listeners, listener)
598 - }
599 - }
681 + _, listeners := e.activeSnapshotLocked()
682 e.mu.RUnlock()
683
684 addrs := make([]string, 0, len(listeners))
@@ -638,21 +720,19 @@ func (e *Exposure) WaitDatagramReady(ctx context.Context) ([]string, error) {
720 func (e *Exposure) monitorStartupCounts() {
721 ticker := time.NewTicker(time.Second)
722 defer ticker.Stop()
641 - prevStatuses := make(map[string]listenerStatus)
642 - firstRun := true
723 + var (
724 + lastStatuses map[string]listenerStatus
725 + lastBannedCount = -1
726 + )
727
728 for {
729 e.mu.RLock()
646 - listeners := make([]*Listener, 0, len(e.listeners))
647 - for _, relayURL := range e.activeRelayURLs {
648 - if listener, ok := e.listeners[relayURL]; ok {
649 - listeners = append(listeners, listener)
650 - }
651 - }
730 + _, listeners := e.activeSnapshotLocked()
731 bannedCount := len(e.bannedRelayURLs)
732 e.mu.RUnlock()
733
655 - readyCount, inactiveCount := 0, 0
734 + currentStatuses := make(map[string]listenerStatus, len(listeners))
735 + readyCount := 0
736 activated := make([]string, 0)
737 deactivated := make([]string, 0)
738
@@ -663,23 +743,40 @@ func (e *Exposure) monitorStartupCounts() {
743 relayURL := listener.api.baseURL.String()
744
745 status := listener.StartupStatus()
746 + currentStatuses[relayURL] = status
747 if status == listenerStatusReady {
748 readyCount++
668 - } else {
669 - inactiveCount++
749 }
750
672 - if prev, ok := prevStatuses[relayURL]; ok && prev != status {
673 - if status == listenerStatusReady {
674 - activated = append(activated, relayURL)
675 - } else {
676 - deactivated = append(deactivated, relayURL)
751 + if lastStatuses != nil {
752 + if prev, ok := lastStatuses[relayURL]; ok && prev != status {
753 + if status == listenerStatusReady {
754 + activated = append(activated, relayURL)
755 + } else {
756 + deactivated = append(deactivated, relayURL)
757 + }
758 + }
759 + }
760 + }
761 +
762 + changed := lastStatuses == nil || bannedCount != lastBannedCount || len(currentStatuses) != len(lastStatuses)
763 + if !changed {
764 + for relayURL, status := range currentStatuses {
765 + if lastStatuses[relayURL] != status {
766 + changed = true
767 + break
768 }
769 }
679 - prevStatuses[relayURL] = status
770 }
771 + if changed {
772 + inactiveCount := len(currentStatuses) - readyCount
773 + for relayURL, status := range lastStatuses {
774 + if _, ok := currentStatuses[relayURL]; ok || status != listenerStatusReady {
775 + continue
776 + }
777 + deactivated = append(deactivated, relayURL)
778 + }
779
682 - if firstRun || len(activated) > 0 || len(deactivated) > 0 {
780 event := log.Info().
781 Int("banned", bannedCount).
782 Int("inactive", inactiveCount).
@@ -691,7 +788,8 @@ func (e *Exposure) monitorStartupCounts() {
788 event = event.Strs("deactivated", deactivated)
789 }
790 event.Msg("relay status")
694 - firstRun = false
791 + lastStatuses = currentStatuses
792 + lastBannedCount = bannedCount
793 }
794
795 select {
@@ -711,70 +809,96 @@ func (e *Exposure) runDiscoveryLoop(ctx context.Context) {
809
810 for {
811 e.mu.RLock()
714 - knownRelayURLs := append([]string(nil), e.knownRelayURLs...)
812 + relayURLs := append([]string(nil), e.knownRelayURLs...)
813 e.mu.RUnlock()
814 + if len(relayURLs) == 0 {
815 + select {
816 + case <-ctx.Done():
817 + return
818 + case <-ticker.C:
819 + continue
820 + }
821 + }
822
717 - peers, err := utils.ExcludeLocalRelayURLs(knownRelayURLs...)
718 - if err == nil && len(peers) > 0 {
719 - relayURLs := append([]string(nil), peers...)
720 - var discoverErr error
721 -
722 - for _, peer := range peers {
723 - resp, err := discovery.Discover(ctx, peer, types.DiscoverRequest{}, e.rootCAPEM, nil)
724 - if err != nil {
725 - discoverErr = errors.Join(discoverErr, fmt.Errorf("discover %q: %w", peer, err))
726 - continue
727 - }
823 + discoveredRelayURLs := append([]string(nil), relayURLs...)
824 + successCount := 0
825 + var discoveryErr error
826 + var warnErr error
827
729 - now := time.Now().UTC()
730 - self, descriptors, err := discovery.ValidateResponse(resp, now)
731 - if err != nil {
732 - if self.RelayID == "" {
733 - discoverErr = errors.Join(discoverErr, fmt.Errorf("validate %q self descriptor: %w", peer, err))
734 - continue
735 - }
736 - discoverErr = errors.Join(discoverErr, fmt.Errorf("validate %q peer descriptors: %w", peer, err))
828 + for _, relayURL := range relayURLs {
829 + resp, err := discovery.Discover(ctx, relayURL, types.DiscoverRequest{}, e.rootCAPEM, nil)
830 + if err != nil {
831 + if ctx.Err() != nil {
832 + return
833 }
834 + discoveryErr = errors.Join(discoveryErr, fmt.Errorf("discover %q: %w", relayURL, err))
835 + continue
836 + }
837
739 - urls := make([]string, 0, 1+len(descriptors))
740 - if apiURL := strings.TrimSpace(self.APIHTTPSAddr); apiURL != "" {
741 - urls = append(urls, apiURL)
742 - }
743 - for _, descriptor := range descriptors {
744 - if apiURL := strings.TrimSpace(descriptor.APIHTTPSAddr); apiURL != "" {
745 - urls = append(urls, apiURL)
746 - }
747 - }
838 + selfDescriptor, peerDescriptors, validateErr := discovery.ValidateResponse(resp, time.Now().UTC())
839 + if selfDescriptor.RelayID == "" || strings.TrimSpace(selfDescriptor.APIHTTPSAddr) == "" {
840 + discoveryErr = errors.Join(discoveryErr, fmt.Errorf("discover %q: missing self descriptor", relayURL))
841 + continue
842 + }
843 + if err := e.pinDiscoverySelfDescriptor(relayURL, selfDescriptor); err != nil {
844 + discoveryErr = errors.Join(discoveryErr, fmt.Errorf("discover %q: %w", relayURL, err))
845 + continue
846 + }
847 + if validateErr != nil {
848 + warnErr = errors.Join(warnErr, fmt.Errorf("discover %q: %w", relayURL, validateErr))
849 + }
850
749 - discoveredRelayURLs, err := utils.ExcludeLocalRelayURLs(urls...)
750 - if err != nil {
751 - discoverErr = errors.Join(discoverErr, fmt.Errorf("extract %q relay urls: %w", peer, err))
851 + descriptorRelayURLs := make([]string, 0, 1+len(peerDescriptors))
852 + descriptorRelayURLs = append(descriptorRelayURLs, selfDescriptor.APIHTTPSAddr)
853 + for _, descriptor := range peerDescriptors {
854 + if err := e.pinDiscoveredDescriptor(descriptor); err != nil {
855 + warnErr = errors.Join(warnErr, fmt.Errorf("discover %q peer %q: %w", relayURL, descriptor.RelayID, err))
856 continue
857 }
754 - relayURLs, err = utils.MergeRelayURLs(relayURLs, nil, discoveredRelayURLs)
755 - if err != nil {
756 - discoverErr = errors.Join(discoverErr, fmt.Errorf("merge %q relay urls: %w", peer, err))
757 - continue
858 + if apiURL := strings.TrimSpace(descriptor.APIHTTPSAddr); apiURL != "" {
859 + descriptorRelayURLs = append(descriptorRelayURLs, apiURL)
860 }
861 }
862
761 - err = discoverErr
762 - switch {
763 - case err == nil:
764 - recovered := discoveryFailed
765 - discoveryFailed = false
766 - added, err := e.setRelayURLs(relayURLs, false)
767 - if err != nil {
768 - log.Warn().
769 - Err(err).
770 - Int("relay_count", len(peers)).
771 - Msg("discover relay urls failed")
772 - } else if recovered || len(added) > 0 {
773 - e.mu.RLock()
774 - totalKnownRelayCount := len(e.knownRelayURLs)
775 - e.mu.RUnlock()
776 - event := log.Info().
777 - Int("peer_count", len(peers)).
863 + discoveredRelayURLs, err = utils.MergeRelayURLs(discoveredRelayURLs, nil, descriptorRelayURLs)
864 + if err != nil {
865 + discoveryErr = errors.Join(discoveryErr, fmt.Errorf("merge %q discovery relays: %w", relayURL, err))
866 + continue
867 + }
868 + successCount++
869 + }
870 +
871 + switch {
872 + case successCount > 0:
873 + recovered := discoveryFailed
874 + discoveryFailed = false
875 + added, err := e.setRelayURLs(discoveredRelayURLs, false)
876 + if err != nil {
877 + log.Warn().
878 + Err(err).
879 + Int("relay_count", len(relayURLs)).
880 + Msg("relay discovery update failed")
881 + } else if recovered || len(added) > 0 || warnErr != nil || discoveryErr != nil {
882 + e.mu.RLock()
883 + totalKnownRelayCount := len(e.knownRelayURLs)
884 + e.mu.RUnlock()
885 + logErr := errors.Join(warnErr, discoveryErr)
886 + event := log.Info().
887 + Int("relay_count", len(relayURLs)).
888 + Int("discovered_count", successCount).
889 + Int("total_known_relay_count", totalKnownRelayCount)
890 + if recovered {
891 + event = event.Bool("recovered", true)
892 + }
893 + if len(added) > 0 {
894 + event = event.Int("added_count", len(added)).
895 + Strs("added_relays", added)
896 + }
897 + if logErr != nil {
898 + event = log.Warn().
899 + Err(logErr).
900 + Int("relay_count", len(relayURLs)).
901 + Int("discovered_count", successCount).
902 Int("total_known_relay_count", totalKnownRelayCount)
903 if recovered {
904 event = event.Bool("recovered", true)
@@ -783,27 +907,17 @@ func (e *Exposure) runDiscoveryLoop(ctx context.Context) {
907 event = event.Int("added_count", len(added)).
908 Strs("added_relays", added)
909 }
786 - event.Msg("discovery relays updated")
910 }
788 - case ctx.Err() != nil:
789 - return
790 - default:
791 - if !discoveryFailed {
792 - log.Debug().
793 - Err(err).
794 - Int("relay_count", len(peers)).
795 - Msg("discover relay urls failed")
796 - }
797 - discoveryFailed = true
798 - }
799 - } else if err != nil {
800 - if ctx.Err() != nil {
801 - return
911 + event.Msg("relay discovery updated")
912 }
913 + case ctx.Err() != nil:
914 + return
915 + default:
916 if !discoveryFailed {
917 log.Debug().
805 - Err(err).
806 - Msg("discover relay urls failed")
918 + Err(discoveryErr).
919 + Int("relay_count", len(relayURLs)).
920 + Msg("relay discovery failed")
921 }
922 discoveryFailed = true
923 }
sdk/expose_test.go
+104 -12
@@ -3,6 +3,8 @@ package sdk
3 import (
4 "net/url"
5 "testing"
6 +
7 + "github.com/gosuda/portal/v2/types"
8 )
9
10 func TestExposureBanRelayURLMovesRelay(t *testing.T) {
@@ -23,7 +25,6 @@ func TestExposureBanRelayURLMovesRelay(t *testing.T) {
25
26 exposure := &Exposure{
27 knownRelayURLs: []string{relayA, relayB},
26 - activeRelayURLs: []string{relayA, relayB},
28 bannedRelayURLs: nil,
29 listeners: map[string]*Listener{
30 relayA: listener,
@@ -33,19 +34,21 @@ func TestExposureBanRelayURLMovesRelay(t *testing.T) {
34
35 exposure.banRelayURL(relayA)
36
36 - if got := exposure.KnownRelayURLs(); len(got) != 1 || got[0] != relayB {
37 - t.Fatalf("KnownRelayURLs() = %v, want [%q]", got, relayB)
38 - }
37 if got := exposure.ActiveRelayURLs(); len(got) != 1 || got[0] != relayB {
38 t.Fatalf("ActiveRelayURLs() = %v, want [%q]", got, relayB)
39 }
42 - if got := exposure.BannedRelayURLs(); len(got) != 1 || got[0] != relayA {
43 - t.Fatalf("BannedRelayURLs() = %v, want [%q]", got, relayA)
44 - }
40
41 exposure.mu.RLock()
42 + knownRelayURLs := append([]string(nil), exposure.knownRelayURLs...)
43 + bannedRelayURLs := append([]string(nil), exposure.bannedRelayURLs...)
44 _, listenerExists := exposure.listeners[relayA]
45 exposure.mu.RUnlock()
46 + if len(knownRelayURLs) != 1 || knownRelayURLs[0] != relayB {
47 + t.Fatalf("knownRelayURLs = %v, want [%q]", knownRelayURLs, relayB)
48 + }
49 + if len(bannedRelayURLs) != 1 || bannedRelayURLs[0] != relayA {
50 + t.Fatalf("bannedRelayURLs = %v, want [%q]", bannedRelayURLs, relayA)
51 + }
52 if listenerExists {
53 t.Fatal("banned relay listener still exists in exposure.listeners")
54 }
@@ -71,13 +74,102 @@ func TestExposureSetRelayURLsSkipsBannedRelay(t *testing.T) {
74 if len(added) != 1 || added[0] != relayA {
75 t.Fatalf("added relay urls = %v, want [%q]", added, relayA)
76 }
74 - if got := exposure.KnownRelayURLs(); len(got) != 1 || got[0] != relayA {
75 - t.Fatalf("KnownRelayURLs() = %v, want [%q]", got, relayA)
76 - }
77 if got := exposure.ActiveRelayURLs(); len(got) != 1 || got[0] != relayA {
78 t.Fatalf("ActiveRelayURLs() = %v, want [%q]", got, relayA)
79 }
80 - if got := exposure.BannedRelayURLs(); len(got) != 1 || got[0] != relayB {
81 - t.Fatalf("BannedRelayURLs() = %v, want [%q]", got, relayB)
80 + exposure.mu.RLock()
81 + knownRelayURLs := append([]string(nil), exposure.knownRelayURLs...)
82 + bannedRelayURLs := append([]string(nil), exposure.bannedRelayURLs...)
83 + exposure.mu.RUnlock()
84 + if len(knownRelayURLs) != 1 || knownRelayURLs[0] != relayA {
85 + t.Fatalf("knownRelayURLs = %v, want [%q]", knownRelayURLs, relayA)
86 + }
87 + if len(bannedRelayURLs) != 1 || bannedRelayURLs[0] != relayB {
88 + t.Fatalf("bannedRelayURLs = %v, want [%q]", bannedRelayURLs, relayB)
89 + }
90 +}
91 +
92 +func TestExposureSetRelayURLsRemovesStaleListener(t *testing.T) {
93 + const (
94 + relayA = "https://relay-a.example"
95 + relayB = "https://relay-b.example"
96 + )
97 +
98 + relayAURL, err := url.Parse(relayA)
99 + if err != nil {
100 + t.Fatalf("url.Parse(relayA) error = %v", err)
101 + }
102 + relayBURL, err := url.Parse(relayB)
103 + if err != nil {
104 + t.Fatalf("url.Parse(relayB) error = %v", err)
105 + }
106 +
107 + relayAClosed := make(chan struct{})
108 + exposure := &Exposure{
109 + knownRelayURLs: []string{relayA, relayB},
110 + listeners: map[string]*Listener{
111 + relayA: {
112 + api: &apiClient{baseURL: relayAURL},
113 + cancel: func() { close(relayAClosed) },
114 + doneCh: relayAClosed,
115 + },
116 + relayB: {
117 + api: &apiClient{baseURL: relayBURL},
118 + },
119 + },
120 + }
121 +
122 + added, err := exposure.setRelayURLs([]string{relayB}, false)
123 + if err != nil {
124 + t.Fatalf("setRelayURLs() error = %v", err)
125 + }
126 + if len(added) != 0 {
127 + t.Fatalf("added relay urls = %v, want empty", added)
128 + }
129 +
130 + select {
131 + case <-relayAClosed:
132 + default:
133 + t.Fatal("stale relay listener was not closed")
134 + }
135 +
136 + exposure.mu.RLock()
137 + knownRelayURLs := append([]string(nil), exposure.knownRelayURLs...)
138 + _, relayAExists := exposure.listeners[relayA]
139 + _, relayBExists := exposure.listeners[relayB]
140 + exposure.mu.RUnlock()
141 + if len(knownRelayURLs) != 1 || knownRelayURLs[0] != relayB {
142 + t.Fatalf("knownRelayURLs = %v, want [%q]", knownRelayURLs, relayB)
143 + }
144 + if relayAExists {
145 + t.Fatal("stale relay listener still exists in exposure.listeners")
146 + }
147 + if !relayBExists {
148 + t.Fatal("active relay listener missing from exposure.listeners")
149 + }
150 +}
151 +
152 +func TestExposurePinDiscoveredDescriptorRejectsIdentityChange(t *testing.T) {
153 + exposure := &Exposure{}
154 + desc := types.RelayDescriptor{
155 + RelayID: "relay-a",
156 + APIHTTPSAddr: "https://relay-a.example",
157 + SignerPublicKey: "signer-a",
158 + }
159 +
160 + if err := exposure.pinDiscoverySelfDescriptor(desc.APIHTTPSAddr, desc); err != nil {
161 + t.Fatalf("pinDiscoverySelfDescriptor() error = %v", err)
162 + }
163 +
164 + changedSigner := desc
165 + changedSigner.SignerPublicKey = "signer-b"
166 + if err := exposure.pinDiscoveredDescriptor(changedSigner); err == nil {
167 + t.Fatal("pinDiscoveredDescriptor() error = nil, want pinned signer mismatch")
168 + }
169 +
170 + changedURL := desc
171 + changedURL.APIHTTPSAddr = "https://relay-b.example"
172 + if err := exposure.pinDiscoveredDescriptor(changedURL); err == nil {
173 + t.Fatal("pinDiscoveredDescriptor() error = nil, want pinned relay url mismatch")
174 }
175 }
sdk/listener.go
+20 -36
@@ -4,7 +4,6 @@ import (
4 "context"
5 "crypto/tls"
6 "errors"
7 - "fmt"
7 "io"
8 "net"
9 "net/url"
@@ -36,9 +35,7 @@ type ListenerConfig struct {
35 ReadyTarget int
36 RetryCount int
37 RetryWait time.Duration
39 -
40 - RegisterBootstraps []string
41 - ownerAddress string
38 + ownerAddress string
39 }
40
41 type listenerStatus string
@@ -54,11 +51,10 @@ type Listener struct {
51 cancel context.CancelFunc
52 doneCh <-chan struct{}
53
57 - retryCount int
58 - retryWait time.Duration
59 - leaseTTL time.Duration
60 - renewBefore time.Duration
61 - registerBootstraps []string
54 + retryCount int
55 + retryWait time.Duration
56 + leaseTTL time.Duration
57 + renewBefore time.Duration
58
59 stream *transport.ClientStream
60 datagram *transport.ClientDatagram
@@ -95,25 +91,18 @@ func NewListener(ctx context.Context, relayURL string, cfg ListenerConfig) (*Lis
91 return nil, err
92 }
93
98 - initialBootstraps, err := utils.NormalizeRelayURLs(cfg.RegisterBootstraps...)
99 - if err != nil {
100 - cancel()
101 - return nil, fmt.Errorf("normalize bootstraps: %w", err)
102 - }
103 -
94 l := &Listener{
105 - doneCh: listenerCtx.Done(),
106 - cancel: cancel,
107 - api: api,
108 - registered: make(chan struct{}),
109 - startupStatus: listenerStatusInactive,
110 - retryCount: cfg.RetryCount,
111 - retryWait: retryWait,
112 - leaseTTL: leaseTTL,
113 - renewBefore: renewBefore,
114 - registerBootstraps: initialBootstraps,
115 - metadata: cfg.Metadata.Copy(),
116 - banMITM: cfg.BanMITM,
95 + doneCh: listenerCtx.Done(),
96 + cancel: cancel,
97 + api: api,
98 + registered: make(chan struct{}),
99 + startupStatus: listenerStatusInactive,
100 + retryCount: cfg.RetryCount,
101 + retryWait: retryWait,
102 + leaseTTL: leaseTTL,
103 + renewBefore: renewBefore,
104 + metadata: cfg.Metadata.Copy(),
105 + banMITM: cfg.BanMITM,
106 }
107 l.mitmManager = newMITMManager(listenerCtx, l)
108 l.stream = transport.NewClientStream(readyTarget, handshakeTimeout)
@@ -138,7 +127,7 @@ func (l *Listener) runStartup(ctx context.Context, readyTarget int) {
127 var retries int
128
129 for {
141 - err := l.registerAndConfigure(ctx, l.registerBootstraps)
130 + err := l.registerAndConfigure(ctx)
131 switch {
132 case err == nil:
133 for range readyTarget {
@@ -449,18 +438,14 @@ func (l *Listener) renewLease(ctx context.Context) error {
438
439 requestCtx, cancel = context.WithTimeout(ctx, 10*time.Second)
440 defer cancel()
452 - if err := l.registerAndConfigure(requestCtx, l.registerBootstraps); err != nil {
441 + if err := l.registerAndConfigure(requestCtx); err != nil {
442 return err
443 }
444 return nil
445 }
446
458 -func (l *Listener) registerAndConfigure(ctx context.Context, registerBootstraps []string) error {
459 - if err := l.api.ensureReady(ctx); err != nil {
460 - return err
461 - }
462 -
463 - resp, err := l.api.registerLease(ctx, l.leaseTTL, l.datagram != nil, registerBootstraps)
447 +func (l *Listener) registerAndConfigure(ctx context.Context) error {
448 + resp, err := l.api.registerLease(ctx, l.leaseTTL, l.datagram != nil)
449 if err != nil {
450 return err
451 }
@@ -471,7 +456,6 @@ func (l *Listener) registerAndConfigure(ctx context.Context, registerBootstraps
456 Message: "relay did not enable required udp support",
457 }
458 }
474 -
459 tlsConf, tlsCloser, err := keyless.BuildClientTLSConfig(l.api.baseURL.String(), []string{resp.Hostname})
460 if err != nil {
461 _ = l.api.unregisterLease(context.Background(), resp.LeaseID)
types/api.go
-3
@@ -59,7 +59,6 @@ type RegisterRequest struct {
59 Metadata LeaseMetadata `json:"metadata"`
60 OwnerAddress string `json:"owner_address,omitempty"`
61 TTL int `json:"ttl,omitempty"`
62 - Bootstraps []string `json:"bootstraps,omitempty"`
62 UDPEnabled bool `json:"udp_enabled,omitempty"`
63 ReportedIP string `json:"reported_ip,omitempty"`
64 }
@@ -67,10 +66,8 @@ type RegisterRequest struct {
66 type RegisterResponse struct {
67 ExpiresAt time.Time `json:"expires_at"`
68 LeaseID string `json:"lease_id"`
70 - ConnectURL string `json:"connect_url"`
69 Hostname string `json:"hostname"`
70 Metadata LeaseMetadata `json:"metadata"`
73 - Bootstraps []string `json:"bootstraps,omitempty"`
71 UDPAddr string `json:"udp_addr,omitempty"`
72 UDPEnabled bool `json:"udp_enabled,omitempty"`
73 }
types/lease.go
-1
@@ -30,7 +30,6 @@ type Lease struct {
30 ClientIP string
31 ReportedIP string
32 Hostname string
33 - Bootstraps []string
33 UDPEnabled bool
34 Metadata LeaseMetadata
35 OwnerAddress string