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 == "" {