refact: remove discovery.go and simplfy refresher
Chang committed
Apr 12, 2026 at 14:56 UTC
229e7a7437fd8d2c5119db699baf47a53a619871
9 files changed
+356
-313
portal/api_server.go
+1
-1
@@ -172,7 +172,7 @@ func (s *Server) handleRelayDiscovery(w http.ResponseWriter, r *http.Request) {
172
overlayCIDRs = append([]string(nil), cfg.OverlayCIDRs...)
173
}
174
175
- self, err := discovery.NormalizeDescriptor(types.RelayDescriptor{
175
+ self, err := utils.NormalizeDescriptor(types.RelayDescriptor{
176
Identity: s.identity.Base(),
177
RelayID: s.cfg.PortalURL,
178
OwnerAddress: s.identity.Address,
portal/discovery/discovery.go
deleted
-232
@@ -1,232 +0,0 @@
1
-package discovery
2
-
3
-import (
4
- "context"
5
- "errors"
6
- "fmt"
7
- "net/http"
8
- "net/url"
9
- "strings"
10
- "time"
11
-
12
- "github.com/rs/zerolog/log"
13
-
14
- "github.com/gosuda/portal-tunnel/v2/portal/keyless"
15
- "github.com/gosuda/portal-tunnel/v2/types"
16
- "github.com/gosuda/portal-tunnel/v2/utils"
17
-)
18
-
19
-const (
20
- defaultRequestTimeout = 15 * time.Second
21
- DiscoveryPollInterval = 1 * time.Minute
22
-)
23
-
24
-func NormalizeDescriptor(desc types.RelayDescriptor) (types.RelayDescriptor, error) {
25
- desc.Name = utils.NormalizeHostname(desc.Name)
26
- desc.Address = strings.TrimSpace(desc.Address)
27
- desc.APIHTTPSAddr = strings.TrimSpace(desc.APIHTTPSAddr)
28
- desc.RelayID = strings.TrimSpace(desc.RelayID)
29
- desc.IngressTLSAddr = strings.TrimSpace(desc.IngressTLSAddr)
30
- desc.WireGuardPublicKey = strings.TrimSpace(desc.WireGuardPublicKey)
31
- desc.WireGuardEndpoint = strings.TrimSpace(desc.WireGuardEndpoint)
32
- desc.OverlayIPv4 = strings.TrimSpace(desc.OverlayIPv4)
33
- desc.OverlayCIDRs = utils.NormalizeIPPrefixes(desc.OverlayCIDRs)
34
- desc.OwnerAddress = strings.TrimSpace(desc.OwnerAddress)
35
- desc.SignerPublicKey = strings.TrimSpace(desc.SignerPublicKey)
36
- if !desc.IssuedAt.IsZero() {
37
- desc.IssuedAt = desc.IssuedAt.UTC()
38
- }
39
- if !desc.ExpiresAt.IsZero() {
40
- desc.ExpiresAt = desc.ExpiresAt.UTC()
41
- }
42
-
43
- if desc.APIHTTPSAddr != "" {
44
- normalized, err := utils.NormalizeRelayURL(desc.APIHTTPSAddr)
45
- if err != nil {
46
- return types.RelayDescriptor{}, fmt.Errorf("normalize api https addr: %w", err)
47
- }
48
- desc.APIHTTPSAddr = normalized
49
- }
50
- if desc.RelayID != "" {
51
- normalized, err := utils.NormalizeRelayURL(desc.RelayID)
52
- if err != nil {
53
- return types.RelayDescriptor{}, fmt.Errorf("normalize relay id: %w", err)
54
- }
55
- desc.RelayID = normalized
56
- }
57
- if desc.RelayID == "" {
58
- desc.RelayID = desc.APIHTTPSAddr
59
- }
60
- if desc.Address != "" {
61
- normalized, err := utils.NormalizeEVMAddress(desc.Address)
62
- if err != nil {
63
- return types.RelayDescriptor{}, fmt.Errorf("normalize address: %w", err)
64
- }
65
- desc.Address = normalized
66
- }
67
- if desc.OwnerAddress == "" {
68
- desc.OwnerAddress = desc.Address
69
- }
70
- if desc.OwnerAddress != "" {
71
- normalized, err := utils.NormalizeEVMAddress(desc.OwnerAddress)
72
- if err != nil {
73
- return types.RelayDescriptor{}, fmt.Errorf("normalize owner address: %w", err)
74
- }
75
- desc.OwnerAddress = normalized
76
- }
77
- if desc.SignerPublicKey == "" {
78
- desc.SignerPublicKey = desc.PublicKey
79
- }
80
- return desc, nil
81
-}
82
-
83
-func ValidateDescriptor(desc types.RelayDescriptor, now time.Time) (types.RelayDescriptor, error) {
84
- normalized, err := NormalizeDescriptor(desc)
85
- if err != nil {
86
- return types.RelayDescriptor{}, err
87
- }
88
- if now.IsZero() {
89
- now = time.Now()
90
- }
91
- now = now.UTC()
92
-
93
- switch {
94
- case normalized.Name == "":
95
- return types.RelayDescriptor{}, errors.New("identity.name is required")
96
- case normalized.APIHTTPSAddr == "":
97
- return types.RelayDescriptor{}, errors.New("api_https_addr is required")
98
- case normalized.RelayID == "":
99
- return types.RelayDescriptor{}, errors.New("relay_id is required")
100
- case normalized.APIHTTPSAddr != "" && normalized.RelayID != normalized.APIHTTPSAddr:
101
- return types.RelayDescriptor{}, errors.New("relay_id must match api_https_addr")
102
- case normalized.Sequence == 0:
103
- return types.RelayDescriptor{}, errors.New("sequence is required")
104
- case normalized.Version == 0:
105
- return types.RelayDescriptor{}, errors.New("version is required")
106
- case normalized.IssuedAt.IsZero():
107
- return types.RelayDescriptor{}, errors.New("issued_at is required")
108
- case normalized.ExpiresAt.IsZero():
109
- return types.RelayDescriptor{}, errors.New("expires_at is required")
110
- case normalized.ExpiresAt.Before(now):
111
- return types.RelayDescriptor{}, errors.New("descriptor expired")
112
- case normalized.IssuedAt.After(normalized.ExpiresAt):
113
- return types.RelayDescriptor{}, errors.New("issued_at must be before expires_at")
114
- }
115
- return normalized, nil
116
-}
117
-
118
-func ValidateRelayDiscoveryResponse(resp types.DiscoveryResponse, now time.Time) (types.RelayDescriptor, []types.RelayDescriptor, error) {
119
- protocolVersion := strings.TrimSpace(resp.ProtocolVersion)
120
- if protocolVersion != types.ProtocolVersion {
121
- return types.RelayDescriptor{}, nil, fmt.Errorf("relay protocol version mismatch: relay=%q client=%q", protocolVersion, types.ProtocolVersion)
122
- }
123
-
124
- self, err := ValidateDescriptor(resp.Self, now)
125
- if err != nil {
126
- return types.RelayDescriptor{}, nil, err
127
- }
128
-
129
- seen := map[string]struct{}{self.Key(): {}}
130
- relays := make([]types.RelayDescriptor, 0, len(resp.Relays))
131
- for _, descriptor := range resp.Relays {
132
- verified, err := ValidateDescriptor(descriptor, now)
133
- if err != nil {
134
- log.Warn().
135
- Err(err).
136
- Str("relay", strings.TrimSpace(descriptor.APIHTTPSAddr)).
137
- Str("name", strings.TrimSpace(descriptor.Name)).
138
- Msg("skipping invalid discovery relay hint")
139
- continue
140
- }
141
- identityKey := verified.Key()
142
- if _, ok := seen[identityKey]; ok {
143
- log.Debug().
144
- Str("relay", verified.APIHTTPSAddr).
145
- Str("identity_key", identityKey).
146
- Msg("skipping duplicate discovery relay hint")
147
- continue
148
- }
149
- seen[identityKey] = struct{}{}
150
- relays = append(relays, verified)
151
- }
152
- return self, relays, nil
153
-}
154
-
155
-// ValidateDescriptorTarget checks if a descriptor matches expected target identity.
156
-func ValidateDescriptorTarget(desc types.RelayDescriptor, targetIdentity types.Identity, targetURL string) error {
157
- normalized, err := NormalizeDescriptor(desc)
158
- if err != nil {
159
- return err
160
- }
161
-
162
- targetName := strings.TrimSpace(targetIdentity.Name)
163
- if targetName != "" {
164
- normalizedTargetName := utils.NormalizeHostname(targetName)
165
- if normalized.Name != normalizedTargetName {
166
- return errors.New("descriptor name does not match target relay")
167
- }
168
- }
169
- targetAddress := strings.TrimSpace(targetIdentity.Address)
170
- if targetAddress != "" {
171
- normalizedTargetAddress, err := utils.NormalizeEVMAddress(targetAddress)
172
- if err != nil {
173
- return err
174
- }
175
- if normalized.Address != normalizedTargetAddress {
176
- return errors.New("descriptor address does not match target relay")
177
- }
178
- }
179
-
180
- if targetURL != "" {
181
- normalizedTargetURL, err := utils.NormalizeRelayURL(targetURL)
182
- if err != nil {
183
- return err
184
- }
185
- if normalized.APIHTTPSAddr != normalizedTargetURL {
186
- return errors.New("descriptor api_https_addr does not match target url")
187
- }
188
- }
189
- return nil
190
-}
191
-
192
-func DiscoverRelayDiscovery(ctx context.Context, baseURL string, rootCAPEM []byte, httpClient *http.Client) (types.DiscoveryResponse, error) {
193
- parsedBaseURL, err := url.Parse(baseURL)
194
- if err != nil {
195
- return types.DiscoveryResponse{}, fmt.Errorf("parse discovery base url: %w", err)
196
- }
197
-
198
- client := httpClient
199
- if client == nil {
200
- _, client, err = keyless.NewRelayHTTPClient(ctx, parsedBaseURL, rootCAPEM, defaultRequestTimeout)
201
- if err != nil {
202
- return types.DiscoveryResponse{}, err
203
- }
204
- }
205
- if client.Timeout == 0 {
206
- clone := *client
207
- clone.Timeout = defaultRequestTimeout
208
- client = &clone
209
- }
210
-
211
- var resp types.DiscoveryResponse
212
- if err := utils.HTTPDoAPIPath(ctx, client, parsedBaseURL, http.MethodGet, types.PathDiscovery, nil, nil, &resp); err != nil {
213
- return types.DiscoveryResponse{}, err
214
- }
215
- return resp, nil
216
-}
217
-
218
-func RequireOverlayRelayDescriptor(desc types.RelayDescriptor) error {
219
- if !desc.SupportsOverlayPeer {
220
- return errors.New("descriptor does not support overlay peer")
221
- }
222
- if desc.WireGuardPublicKey == "" {
223
- return errors.New("descriptor wireguard public key is required")
224
- }
225
- if desc.WireGuardEndpoint == "" {
226
- return errors.New("descriptor wireguard endpoint is required")
227
- }
228
- if desc.OverlayIPv4 == "" {
229
- return errors.New("descriptor overlay ipv4 is required")
230
- }
231
- return nil
232
-}
portal/discovery/policy.go
new
+28
@@ -0,0 +1,28 @@
1
+package discovery
2
+
3
+import (
4
+ "errors"
5
+ "time"
6
+
7
+ "github.com/gosuda/portal-tunnel/v2/types"
8
+)
9
+
10
+const (
11
+ DiscoveryPollInterval = 1 * time.Minute
12
+)
13
+
14
+func RequireOverlayRelayDescriptor(desc types.RelayDescriptor) error {
15
+ if !desc.SupportsOverlayPeer {
16
+ return errors.New("descriptor does not support overlay peer")
17
+ }
18
+ if desc.WireGuardPublicKey == "" {
19
+ return errors.New("descriptor wireguard public key is required")
20
+ }
21
+ if desc.WireGuardEndpoint == "" {
22
+ return errors.New("descriptor wireguard endpoint is required")
23
+ }
24
+ if desc.OverlayIPv4 == "" {
25
+ return errors.New("descriptor overlay ipv4 is required")
26
+ }
27
+ return nil
28
+}
portal/discovery/refresher.go
+37
-4
@@ -2,12 +2,17 @@ package discovery
2
3
import (
4
"context"
5
+ "crypto/tls"
6
"errors"
7
+ "fmt"
8
+ "net/http"
9
+ "net/url"
10
"time"
11
12
"github.com/rs/zerolog/log"
13
14
"github.com/gosuda/portal-tunnel/v2/types"
15
+ "github.com/gosuda/portal-tunnel/v2/utils"
16
)
17
18
const (
@@ -21,7 +26,7 @@ type OverlayRuntime interface {
26
27
type Refresher struct {
28
relaySet *RelaySet
24
- rootCAPEM []byte
29
+ httpClient *http.Client
30
overlay OverlayRuntime
31
directRecoveryFailures int
32
overlayRecoveryFailures int
@@ -31,9 +36,24 @@ func NewRefresher(relaySet *RelaySet, rootCAPEM []byte, overlay OverlayRuntime)
36
if relaySet == nil {
37
return nil, errors.New("relay set is required")
38
}
39
+ httpClient := http.DefaultClient
40
+ if len(rootCAPEM) > 0 {
41
+ rootCAs, err := utils.CertPoolFromPEM(rootCAPEM)
42
+ if err != nil {
43
+ return nil, err
44
+ }
45
+ httpClient = &http.Client{
46
+ Transport: &http.Transport{
47
+ TLSClientConfig: &tls.Config{
48
+ MinVersion: tls.VersionTLS12,
49
+ RootCAs: rootCAs,
50
+ },
51
+ },
52
+ }
53
+ }
54
return &Refresher{
55
relaySet: relaySet,
36
- rootCAPEM: append([]byte(nil), rootCAPEM...),
56
+ httpClient: httpClient,
57
overlay: overlay,
58
directRecoveryFailures: defaultRecoveryFailures,
59
overlayRecoveryFailures: defaultRecoveryFailures,
@@ -58,7 +78,7 @@ func (r *Refresher) Refresh(ctx context.Context) error {
78
79
func (r *Refresher) refreshHTTPS(ctx context.Context) error {
80
for _, bootstrap := range r.relaySet.BootstrapDescriptors() {
61
- resp, err := DiscoverRelayDiscovery(ctx, bootstrap.APIHTTPSAddr, r.rootCAPEM, nil)
81
+ resp, err := r.discoverHTTPS(ctx, bootstrap)
82
if err != nil {
83
if ctx.Err() != nil {
84
return ctx.Err()
@@ -80,7 +100,7 @@ func (r *Refresher) refreshHTTPS(ctx context.Context) error {
100
if r.overlay != nil && relay.SupportsOverlayPeer {
101
continue
102
}
83
- resp, err := DiscoverRelayDiscovery(ctx, relay.APIHTTPSAddr, r.rootCAPEM, nil)
103
+ resp, err := r.discoverHTTPS(ctx, relay)
104
if err != nil {
105
if ctx.Err() != nil {
106
return ctx.Err()
@@ -99,6 +119,19 @@ func (r *Refresher) refreshHTTPS(ctx context.Context) error {
119
return ctx.Err()
120
}
121
122
+func (r *Refresher) discoverHTTPS(ctx context.Context, relay types.RelayDescriptor) (types.DiscoveryResponse, error) {
123
+ baseURL, err := url.Parse(relay.APIHTTPSAddr)
124
+ if err != nil {
125
+ return types.DiscoveryResponse{}, fmt.Errorf("parse discovery base url: %w", err)
126
+ }
127
+
128
+ var resp types.DiscoveryResponse
129
+ if err := utils.HTTPDoAPIPath(ctx, r.httpClient, baseURL, http.MethodGet, types.PathDiscovery, nil, nil, &resp); err != nil {
130
+ return types.DiscoveryResponse{}, err
131
+ }
132
+ return resp, nil
133
+}
134
+
135
func (r *Refresher) refreshOverlay(ctx context.Context) error {
136
for _, relay := range r.relaySet.SyncableDescriptors() {
137
var failureErr error
portal/discovery/relayset.go
+224
-71
@@ -2,6 +2,7 @@ package discovery
2
3
import (
4
"errors"
5
+ "fmt"
6
"net/http"
7
"reflect"
8
"slices"
@@ -10,6 +11,8 @@ import (
11
"sync"
12
"time"
13
14
+ "github.com/rs/zerolog/log"
15
+
16
"github.com/gosuda/portal-tunnel/v2/types"
17
"github.com/gosuda/portal-tunnel/v2/utils"
18
)
@@ -34,6 +37,19 @@ type RelayState struct {
37
consecutiveFailures int
38
}
39
40
+// RelayCandidate is a relay fact projected for client-side selection.
41
+// Discovery owns the facts; callers decide how many candidates to use.
42
+type RelayCandidate struct {
43
+ Descriptor types.RelayDescriptor
44
+ Bootstrap bool
45
+ FirstSeenAt time.Time
46
+ LastSeenAt time.Time
47
+}
48
+
49
+type SelectionOptions struct {
50
+ Limit int
51
+}
52
+
53
func (s RelayState) isDefaultLocalState() bool {
54
return !s.Banned && s.status == relayStatusHinted && s.consecutiveFailures == 0
55
}
@@ -209,18 +225,90 @@ func (s *RelaySet) ActiveRelayURLs() []string {
225
s.mu.RLock()
226
defer s.mu.RUnlock()
227
212
- now := time.Now().UTC()
228
+ return selectRelayURLsLocked(s.relayCandidatesLocked(time.Now().UTC()), SelectionOptions{})
229
+}
230
+
231
+func (s *RelaySet) RelayCandidates() []RelayCandidate {
232
+ s.mu.RLock()
233
+ defer s.mu.RUnlock()
234
+
235
+ return s.relayCandidatesLocked(time.Now().UTC())
236
+}
237
+
238
+func (s *RelaySet) SelectRelayURLs(opts SelectionOptions) []string {
239
+ s.mu.RLock()
240
+ defer s.mu.RUnlock()
241
+
242
+ return selectRelayURLsLocked(s.relayCandidatesLocked(time.Now().UTC()), opts)
243
+}
244
+
245
+func selectRelayURLsLocked(candidates []RelayCandidate, opts SelectionOptions) []string {
246
+ if len(candidates) == 0 {
247
+ return nil
248
+ }
249
+
250
+ limit := opts.Limit
251
+ out := make([]string, 0, len(candidates))
252
+ seen := make(map[string]struct{}, len(candidates))
253
+ for _, candidate := range candidates {
254
+ relayURL := relayCandidateURL(candidate)
255
+ if relayURL == "" {
256
+ continue
257
+ }
258
+ if _, ok := seen[relayURL]; ok {
259
+ continue
260
+ }
261
+ seen[relayURL] = struct{}{}
262
+ out = append(out, relayURL)
263
+ if limit > 0 && len(out) >= limit {
264
+ break
265
+ }
266
+ }
267
+ if len(out) == 0 {
268
+ return nil
269
+ }
270
+ return out
271
+}
272
+
273
+func relayCandidateURL(candidate RelayCandidate) string {
274
+ relayURL := strings.TrimSpace(candidate.Descriptor.APIHTTPSAddr)
275
+ if relayURL != "" {
276
+ return relayURL
277
+ }
278
+ return strings.TrimSpace(candidate.Descriptor.RelayID)
279
+}
280
+
281
+func (s *RelaySet) relayCandidatesLocked(now time.Time) []RelayCandidate {
282
bootstrapRelayURLs := s.bootstrapRelayURLsLocked()
283
projections := s.descriptorProjectionsLocked()
284
216
- out := make([]string, 0, len(bootstrapRelayURLs)+len(projections))
285
+ out := make([]RelayCandidate, 0, len(bootstrapRelayURLs)+len(projections))
286
seen := make(map[string]struct{}, len(bootstrapRelayURLs)+len(projections))
287
for _, relayURL := range bootstrapRelayURLs {
288
if _, ok := seen[relayURL]; ok {
289
continue
290
}
291
seen[relayURL] = struct{}{}
223
- out = append(out, relayURL)
292
+
293
+ candidate := RelayCandidate{
294
+ Bootstrap: true,
295
+ Descriptor: types.RelayDescriptor{
296
+ Identity: types.Identity{
297
+ Name: utils.PortalRootHost(relayURL),
298
+ },
299
+ RelayID: relayURL,
300
+ APIHTTPSAddr: relayURL,
301
+ Version: 1,
302
+ },
303
+ }
304
+ if relayKey, ok := s.relayKeysByURL[relayURL]; ok {
305
+ if record, ok := s.relays[relayKey]; ok && record.Descriptor.APIHTTPSAddr != "" {
306
+ candidate.Descriptor = record.Descriptor
307
+ candidate.FirstSeenAt = record.FirstSeenAt
308
+ candidate.LastSeenAt = record.LastSeenAt
309
+ }
310
+ }
311
+ out = append(out, candidate)
312
}
313
for _, projection := range projections {
314
if projection.state.Banned || projection.state.status != relayStatusConfirmed || relayExpiredAt(projection.state, now) {
@@ -230,7 +318,12 @@ func (s *RelaySet) ActiveRelayURLs() []string {
318
continue
319
}
320
seen[projection.relayURL] = struct{}{}
233
- out = append(out, projection.relayURL)
321
+ out = append(out, RelayCandidate{
322
+ Descriptor: projection.state.Descriptor,
323
+ Bootstrap: projection.bootstrap,
324
+ FirstSeenAt: projection.state.FirstSeenAt,
325
+ LastSeenAt: projection.state.LastSeenAt,
326
+ })
327
}
328
if len(out) == 0 {
329
return nil
@@ -415,16 +508,12 @@ func (s *RelaySet) SetBootstrapRelayURLs(inputs []string) error {
508
return nil
509
}
510
418
-func (s *RelaySet) registerDescriptor(desc types.RelayDescriptor, now time.Time) (string, bool, bool, error) {
419
- normalized, err := NormalizeDescriptor(desc)
420
- if err != nil {
421
- return "", false, false, err
422
- }
423
- relayKey := normalized.Key()
511
+func (s *RelaySet) storeDescriptorLocked(desc types.RelayDescriptor, now time.Time) (string, bool, bool, error) {
512
+ relayKey := desc.Key()
513
if relayKey == "" {
514
return "", false, false, errors.New("descriptor identity is required")
515
}
427
- if knownRelayKey, ok := s.relayKeysByURL[normalized.APIHTTPSAddr]; ok && knownRelayKey != relayKey {
516
+ if knownRelayKey, ok := s.relayKeysByURL[desc.APIHTTPSAddr]; ok && knownRelayKey != relayKey {
517
return "", false, false, errors.New("descriptor identity does not match known relay url")
518
}
519
if now.IsZero() {
@@ -438,11 +527,11 @@ func (s *RelaySet) registerDescriptor(desc types.RelayDescriptor, now time.Time)
527
}
528
previousURL := record.Descriptor.APIHTTPSAddr
529
previousDescriptor := record.Descriptor
441
- record.Descriptor = normalized
530
+ record.Descriptor = desc
531
record.LastSeenAt = now
532
s.relays[relayKey] = record
444
- s.relayKeysByURL[normalized.APIHTTPSAddr] = relayKey
445
- if previousURL != "" && previousURL != normalized.APIHTTPSAddr {
533
+ s.relayKeysByURL[desc.APIHTTPSAddr] = relayKey
534
+ if previousURL != "" && previousURL != desc.APIHTTPSAddr {
535
delete(s.relayKeysByURL, previousURL)
536
state := s.localByURL[previousURL]
537
bootstrap := false
@@ -456,27 +545,135 @@ func (s *RelaySet) registerDescriptor(desc types.RelayDescriptor, now time.Time)
545
}
546
}
547
459
- changed := added || !reflect.DeepEqual(previousDescriptor, normalized)
548
+ changed := added || !reflect.DeepEqual(previousDescriptor, desc)
549
return relayKey, added, changed, nil
550
}
551
463
-func (s *RelaySet) applyDiscoveryDescriptorsLocked(targetIdentity types.Identity, targetURL string, selfDescriptor types.RelayDescriptor, relayDescriptors []types.RelayDescriptor, now time.Time) (relaySetChanged bool, err error) {
464
- if strings.TrimSpace(targetIdentity.Name) == "" && strings.TrimSpace(targetIdentity.Address) == "" {
465
- return false, errors.New("target relay identity is required")
466
- }
552
+func (s *RelaySet) applyDiscoveryResponse(targetIdentity types.Identity, targetURL string, resp types.DiscoveryResponse, now time.Time, requireOverlay bool) (relaySetChanged bool, warnErr error, err error) {
553
if now.IsZero() {
554
now = time.Now().UTC()
555
+ } else {
556
+ now = now.UTC()
557
+ }
558
+
559
+ protocolVersion := strings.TrimSpace(resp.ProtocolVersion)
560
+ if protocolVersion != types.ProtocolVersion {
561
+ err := fmt.Errorf("relay protocol version mismatch: relay=%q client=%q", protocolVersion, types.ProtocolVersion)
562
+ return false, err, err
563
+ }
564
+
565
+ normalize := func(desc types.RelayDescriptor) (types.RelayDescriptor, error) {
566
+ normalized, err := utils.NormalizeDescriptor(desc)
567
+ if err != nil {
568
+ return types.RelayDescriptor{}, err
569
+ }
570
+
571
+ switch {
572
+ case normalized.Name == "":
573
+ return types.RelayDescriptor{}, errors.New("identity.name is required")
574
+ case normalized.APIHTTPSAddr == "":
575
+ return types.RelayDescriptor{}, errors.New("api_https_addr is required")
576
+ case normalized.RelayID == "":
577
+ return types.RelayDescriptor{}, errors.New("relay_id is required")
578
+ case normalized.APIHTTPSAddr != "" && normalized.RelayID != normalized.APIHTTPSAddr:
579
+ return types.RelayDescriptor{}, errors.New("relay_id must match api_https_addr")
580
+ case normalized.Sequence == 0:
581
+ return types.RelayDescriptor{}, errors.New("sequence is required")
582
+ case normalized.Version == 0:
583
+ return types.RelayDescriptor{}, errors.New("version is required")
584
+ case normalized.IssuedAt.IsZero():
585
+ return types.RelayDescriptor{}, errors.New("issued_at is required")
586
+ case normalized.ExpiresAt.IsZero():
587
+ return types.RelayDescriptor{}, errors.New("expires_at is required")
588
+ case normalized.ExpiresAt.Before(now):
589
+ return types.RelayDescriptor{}, errors.New("descriptor expired")
590
+ case normalized.IssuedAt.After(normalized.ExpiresAt):
591
+ return types.RelayDescriptor{}, errors.New("issued_at must be before expires_at")
592
+ }
593
+ return normalized, nil
594
+ }
595
+
596
+ selfDescriptor, err := normalize(resp.Self)
597
+ if err != nil {
598
+ return false, err, err
599
+ }
600
+ if requireOverlay {
601
+ if err := RequireOverlayRelayDescriptor(selfDescriptor); err != nil {
602
+ return false, warnErr, err
603
+ }
604
+ }
605
+
606
+ seen := map[string]struct{}{selfDescriptor.Key(): {}}
607
+ relayDescriptors := make([]types.RelayDescriptor, 0, len(resp.Relays))
608
+ for _, descriptor := range resp.Relays {
609
+ relayDescriptor, err := normalize(descriptor)
610
+ if err != nil {
611
+ log.Warn().
612
+ Err(err).
613
+ Str("relay", strings.TrimSpace(descriptor.APIHTTPSAddr)).
614
+ Str("name", strings.TrimSpace(descriptor.Name)).
615
+ Msg("skipping invalid discovery relay hint")
616
+ continue
617
+ }
618
+ relayKey := relayDescriptor.Key()
619
+ if _, ok := seen[relayKey]; ok {
620
+ log.Debug().
621
+ Str("relay", relayDescriptor.APIHTTPSAddr).
622
+ Str("identity_key", relayKey).
623
+ Msg("skipping duplicate discovery relay hint")
624
+ continue
625
+ }
626
+ seen[relayKey] = struct{}{}
627
+ relayDescriptors = append(relayDescriptors, relayDescriptor)
628
+ }
629
+
630
+ if strings.TrimSpace(targetIdentity.Name) == "" && strings.TrimSpace(targetIdentity.Address) == "" {
631
+ return false, warnErr, errors.New("target relay identity is required")
632
}
470
- if err := ValidateDescriptorTarget(selfDescriptor, targetIdentity, targetURL); err != nil {
471
- return false, err
633
+ targetName := strings.TrimSpace(targetIdentity.Name)
634
+ if targetName != "" {
635
+ normalizedTargetName := utils.NormalizeHostname(targetName)
636
+ if selfDescriptor.Name != normalizedTargetName {
637
+ return false, warnErr, errors.New("descriptor name does not match target relay")
638
+ }
639
}
640
+ targetAddress := strings.TrimSpace(targetIdentity.Address)
641
+ if targetAddress != "" {
642
+ normalizedTargetAddress, err := utils.NormalizeEVMAddress(targetAddress)
643
+ if err != nil {
644
+ return false, warnErr, err
645
+ }
646
+ if selfDescriptor.Address != normalizedTargetAddress {
647
+ return false, warnErr, errors.New("descriptor address does not match target relay")
648
+ }
649
+ }
650
+ if targetURL != "" {
651
+ normalizedTargetURL, err := utils.NormalizeRelayURL(targetURL)
652
+ if err != nil {
653
+ return false, warnErr, err
654
+ }
655
+ if selfDescriptor.APIHTTPSAddr != normalizedTargetURL {
656
+ return false, warnErr, errors.New("descriptor api_https_addr does not match target url")
657
+ }
658
+ }
659
+
660
+ s.mu.Lock()
661
+ defer s.mu.Unlock()
662
663
apply := func(desc types.RelayDescriptor, advertise bool) error {
664
if !advertise && s.isSelfRelayDescriptorLocked(desc) {
665
return nil
666
}
667
+ if !advertise && requireOverlay {
668
+ if err := RequireOverlayRelayDescriptor(desc); err != nil {
669
+ if warnErr == nil {
670
+ warnErr = err
671
+ }
672
+ return nil
673
+ }
674
+ }
675
479
- _, added, descriptorChanged, err := s.registerDescriptor(desc, now)
676
+ _, added, descriptorChanged, err := s.storeDescriptorLocked(desc, now)
677
if err != nil {
678
return err
679
}
@@ -500,66 +697,22 @@ func (s *RelaySet) applyDiscoveryDescriptorsLocked(targetIdentity types.Identity
697
}
698
699
if err := apply(selfDescriptor, true); err != nil {
503
- return false, err
700
+ return false, warnErr, err
701
}
702
for _, relayDescriptor := range relayDescriptors {
703
if err := apply(relayDescriptor, false); err != nil {
507
- return false, err
704
+ return false, warnErr, err
705
}
706
}
510
- state := s.localByURL[selfDescriptor.APIHTTPSAddr]
511
- state.status = relayStatusConfirmed
512
- state.consecutiveFailures = 0
513
- s.storeLocalStateLocked(selfDescriptor.APIHTTPSAddr, state)
514
- return relaySetChanged, nil
707
+ return relaySetChanged, warnErr, nil
708
}
709
710
func (s *RelaySet) ApplyRelayDiscoveryResponse(targetIdentity types.Identity, targetURL string, resp types.DiscoveryResponse, now time.Time) (relaySetChanged bool, warnErr error, err error) {
518
- selfDescriptor, relayDescriptors, validateErr := ValidateRelayDiscoveryResponse(resp, now)
519
- warnErr = validateErr
520
- if selfDescriptor.Key() == "" {
521
- return false, warnErr, validateErr
522
- }
523
- s.mu.Lock()
524
- relaySetChanged, err = s.applyDiscoveryDescriptorsLocked(targetIdentity, targetURL, selfDescriptor, relayDescriptors, now)
525
- s.mu.Unlock()
526
- if err != nil {
527
- return false, warnErr, err
528
- }
529
- return relaySetChanged, warnErr, nil
711
+ return s.applyDiscoveryResponse(targetIdentity, targetURL, resp, now, false)
712
}
713
714
func (s *RelaySet) ApplyOverlayRelayDiscoveryResponse(targetIdentity types.Identity, targetURL string, resp types.DiscoveryResponse, now time.Time) (relaySetChanged bool, warnErr error, err error) {
533
- selfDescriptor, relayDescriptors, validateErr := ValidateRelayDiscoveryResponse(resp, now)
534
- warnErr = validateErr
535
- if selfDescriptor.Key() == "" {
536
- return false, warnErr, validateErr
537
- }
538
- if err := RequireOverlayRelayDescriptor(selfDescriptor); err != nil {
539
- return false, warnErr, err
540
- }
541
-
542
- filteredRelayDescriptors := make([]types.RelayDescriptor, 0, len(relayDescriptors))
543
- for _, relayDescriptor := range relayDescriptors {
544
- if s.isSelfRelayDescriptorLocked(relayDescriptor) {
545
- continue
546
- }
547
- if err := RequireOverlayRelayDescriptor(relayDescriptor); err != nil {
548
- if warnErr == nil {
549
- warnErr = err
550
- }
551
- continue
552
- }
553
- filteredRelayDescriptors = append(filteredRelayDescriptors, relayDescriptor)
554
- }
555
-
556
- s.mu.Lock()
557
- relaySetChanged, err = s.applyDiscoveryDescriptorsLocked(targetIdentity, targetURL, selfDescriptor, filteredRelayDescriptors, now)
558
- s.mu.Unlock()
559
- if err != nil {
560
- return false, warnErr, err
561
- }
562
- return relaySetChanged, warnErr, nil
715
+ return s.applyDiscoveryResponse(targetIdentity, targetURL, resp, now, true)
716
}
717
718
func (s *RelaySet) RecordDiscoveryFailure(identity types.Identity, relayURL string, err error, recoveryFailures int) (expired bool, expireReason string, consecutiveFailures int) {
portal/server_test.go
+2
-2
@@ -32,7 +32,7 @@ func mustRelayDescriptor(t *testing.T, relayURL string) types.RelayDescriptor {
32
t.Helper()
33
34
now := time.Now().UTC()
35
- desc, err := discovery.NormalizeDescriptor(types.RelayDescriptor{
35
+ desc, err := utils.NormalizeDescriptor(types.RelayDescriptor{
36
Identity: types.Identity{
37
Name: utils.PortalRootHost(relayURL),
38
},
@@ -540,7 +540,7 @@ func TestServerDiscoverySkipsSelfRelayHint(t *testing.T) {
540
541
now := time.Now().UTC()
542
bootstrapDesc := mustRelayDescriptor(t, "https://bootstrap.example.com")
543
- selfHint, err := discovery.NormalizeDescriptor(types.RelayDescriptor{
543
+ selfHint, err := utils.NormalizeDescriptor(types.RelayDescriptor{
544
Identity: server.identity.Base(),
545
RelayID: "https://self-mirror.example.com",
546
Sequence: uint64(now.UnixMilli()),
sdk/expose.go
+3
-2
@@ -254,9 +254,10 @@ func (e *Exposure) WaitDatagramReady(ctx context.Context) ([]string, error) {
254
255
func (e *Exposure) RunHTTP(ctx context.Context, handler http.Handler, localAddr string) error {
256
var relayListener net.Listener
257
+ activeRelayURLs := e.relaySet.ActiveRelayURLs()
258
e.listenerMu.RLock()
259
activeListeners := make([]*Listener, 0, len(e.relayListeners))
259
- for _, relayURL := range e.relaySet.ActiveRelayURLs() {
260
+ for _, relayURL := range activeRelayURLs {
261
listener, ok := e.relayListeners[relayURL]
262
if !ok {
263
continue
@@ -364,8 +365,8 @@ func (e *Exposure) Close() error {
365
}
366
367
func (e *Exposure) reconcileRelayListeners(failOnError bool) error {
367
- e.listenerMu.Lock()
368
activeRelayURLs := e.relaySet.ActiveRelayURLs()
369
+ e.listenerMu.Lock()
370
currentRelayURLs := make([]string, 0, len(e.relayListeners))
371
for relayURL := range e.relayListeners {
372
currentRelayURLs = append(currentRelayURLs, relayURL)
sdk/expose_test.go
+2
-1
@@ -7,6 +7,7 @@ import (
7
8
"github.com/gosuda/portal-tunnel/v2/portal/discovery"
9
"github.com/gosuda/portal-tunnel/v2/types"
10
+ "github.com/gosuda/portal-tunnel/v2/utils"
11
)
12
13
func mustRelaySet(t *testing.T, relayURLs ...string) *discovery.RelaySet {
@@ -23,7 +24,7 @@ func mustRelayDescriptor(t *testing.T, relayName, relayURL string) types.RelayDe
24
t.Helper()
25
26
now := time.Now().UTC()
26
- desc, err := discovery.NormalizeDescriptor(types.RelayDescriptor{
27
+ desc, err := utils.NormalizeDescriptor(types.RelayDescriptor{
28
Identity: types.Identity{
29
Name: relayName,
30
},
utils/identity.go
+59
@@ -30,6 +30,65 @@ func NormalizeIdentity(identity types.Identity) (types.Identity, error) {
30
return normalized, nil
31
}
32
33
+func NormalizeDescriptor(desc types.RelayDescriptor) (types.RelayDescriptor, error) {
34
+ desc.Name = NormalizeHostname(desc.Name)
35
+ desc.Address = strings.TrimSpace(desc.Address)
36
+ desc.APIHTTPSAddr = strings.TrimSpace(desc.APIHTTPSAddr)
37
+ desc.RelayID = strings.TrimSpace(desc.RelayID)
38
+ desc.IngressTLSAddr = strings.TrimSpace(desc.IngressTLSAddr)
39
+ desc.WireGuardPublicKey = strings.TrimSpace(desc.WireGuardPublicKey)
40
+ desc.WireGuardEndpoint = strings.TrimSpace(desc.WireGuardEndpoint)
41
+ desc.OverlayIPv4 = strings.TrimSpace(desc.OverlayIPv4)
42
+ desc.OverlayCIDRs = NormalizeIPPrefixes(desc.OverlayCIDRs)
43
+ desc.OwnerAddress = strings.TrimSpace(desc.OwnerAddress)
44
+ desc.SignerPublicKey = strings.TrimSpace(desc.SignerPublicKey)
45
+ if !desc.IssuedAt.IsZero() {
46
+ desc.IssuedAt = desc.IssuedAt.UTC()
47
+ }
48
+ if !desc.ExpiresAt.IsZero() {
49
+ desc.ExpiresAt = desc.ExpiresAt.UTC()
50
+ }
51
+
52
+ if desc.APIHTTPSAddr != "" {
53
+ normalized, err := NormalizeRelayURL(desc.APIHTTPSAddr)
54
+ if err != nil {
55
+ return types.RelayDescriptor{}, fmt.Errorf("normalize api https addr: %w", err)
56
+ }
57
+ desc.APIHTTPSAddr = normalized
58
+ }
59
+ if desc.RelayID != "" {
60
+ normalized, err := NormalizeRelayURL(desc.RelayID)
61
+ if err != nil {
62
+ return types.RelayDescriptor{}, fmt.Errorf("normalize relay id: %w", err)
63
+ }
64
+ desc.RelayID = normalized
65
+ }
66
+ if desc.RelayID == "" {
67
+ desc.RelayID = desc.APIHTTPSAddr
68
+ }
69
+ if desc.Address != "" {
70
+ normalized, err := NormalizeEVMAddress(desc.Address)
71
+ if err != nil {
72
+ return types.RelayDescriptor{}, fmt.Errorf("normalize address: %w", err)
73
+ }
74
+ desc.Address = normalized
75
+ }
76
+ if desc.OwnerAddress == "" {
77
+ desc.OwnerAddress = desc.Address
78
+ }
79
+ if desc.OwnerAddress != "" {
80
+ normalized, err := NormalizeEVMAddress(desc.OwnerAddress)
81
+ if err != nil {
82
+ return types.RelayDescriptor{}, fmt.Errorf("normalize owner address: %w", err)
83
+ }
84
+ desc.OwnerAddress = normalized
85
+ }
86
+ if desc.SignerPublicKey == "" {
87
+ desc.SignerPublicKey = desc.PublicKey
88
+ }
89
+ return desc, nil
90
+}
91
+
92
func ResolveRelayStateDir(path string) string {
93
trimmed := strings.TrimSpace(path)
94
if trimmed == "" {