refact: WIP separate policy from relayset
Chang committed
Apr 12, 2026 at 16:32 UTC
b211765c33af5bbeff035e18cc1aabed14e87d9f
9 files changed
+513
-673
portal/discovery/policy.go
+267
-10
@@ -2,27 +2,284 @@ package discovery
2
3
import (
4
"errors"
5
+ "fmt"
6
+ "net/http"
7
+ "sort"
8
+ "strings"
9
"time"
10
11
"github.com/gosuda/portal-tunnel/v2/types"
12
+ "github.com/gosuda/portal-tunnel/v2/utils"
13
)
14
15
+type RelayStatus uint8
16
+
17
const (
11
- DiscoveryPollInterval = 1 * time.Minute
18
+ RelayStatusUnknown RelayStatus = iota
19
+ RelayStatusHinted
20
+ RelayStatusConfirmed
21
+ RelayStatusInvalid
22
+ RelayStatusExpired
23
+ RelayStatusBanned
24
)
25
14
-func RequireOverlayRelayDescriptor(desc types.RelayDescriptor) error {
15
- if !desc.SupportsOverlayPeer {
16
- return errors.New("descriptor does not support overlay peer")
26
+type RelayState struct {
27
+ Descriptor types.RelayDescriptor
28
+ Bootstrap bool
29
+ FirstSeenAt time.Time
30
+ LastSeenAt time.Time
31
+ Status RelayStatus
32
+ Active bool
33
+ Advertise bool
34
+ BootstrapDiscovery bool
35
+ DirectDiscovery bool
36
+ OverlayDiscovery bool
37
+ OverlayPeer bool
38
+
39
+ consecutiveFailures int
40
+}
41
+
42
+func newRelayState(desc types.RelayDescriptor, seenAt time.Time) (RelayState, error) {
43
+ state := RelayState{
44
+ Descriptor: desc,
45
+ }
46
+ if seenAt.IsZero() {
47
+ return state, nil
48
+ }
49
+
50
+ seenAt = seenAt.UTC()
51
+ normalized, err := utils.NormalizeDescriptor(desc)
52
+ if err != nil {
53
+ return RelayState{}, err
54
}
18
- if desc.WireGuardPublicKey == "" {
19
- return errors.New("descriptor wireguard public key is required")
55
+
56
+ switch {
57
+ case normalized.Name == "":
58
+ return RelayState{}, errors.New("identity.name is required")
59
+ case normalized.APIHTTPSAddr == "":
60
+ return RelayState{}, errors.New("api_https_addr is required")
61
+ case normalized.RelayID == "":
62
+ return RelayState{}, errors.New("relay_id is required")
63
+ case normalized.RelayID != normalized.APIHTTPSAddr:
64
+ return RelayState{}, errors.New("relay_id must match api_https_addr")
65
+ case normalized.Sequence == 0:
66
+ return RelayState{}, errors.New("sequence is required")
67
+ case normalized.Version == 0:
68
+ return RelayState{}, errors.New("version is required")
69
+ case normalized.IssuedAt.IsZero():
70
+ return RelayState{}, errors.New("issued_at is required")
71
+ case normalized.ExpiresAt.IsZero():
72
+ return RelayState{}, errors.New("expires_at is required")
73
+ case normalized.ExpiresAt.Before(seenAt):
74
+ return RelayState{}, errors.New("descriptor expired")
75
+ case normalized.IssuedAt.After(normalized.ExpiresAt):
76
+ return RelayState{}, errors.New("issued_at must be before expires_at")
77
}
21
- if desc.WireGuardEndpoint == "" {
22
- return errors.New("descriptor wireguard endpoint is required")
78
+
79
+ state.Descriptor = normalized
80
+ state.FirstSeenAt = seenAt
81
+ state.LastSeenAt = seenAt
82
+ return state, nil
83
+}
84
+
85
+func newRelayHintState(relayURL string) RelayState {
86
+ state, _ := newRelayState(
87
+ types.RelayDescriptor{
88
+ Identity: types.Identity{
89
+ Name: utils.PortalRootHost(relayURL),
90
+ },
91
+ RelayID: relayURL,
92
+ APIHTTPSAddr: relayURL,
93
+ Version: 1,
94
+ },
95
+ time.Time{},
96
+ )
97
+ return state
98
+}
99
+
100
+func (DefaultRelayPolicy) DiscoveryStates(targetIdentity types.Identity, targetURL string, selfRelayKey string, selfRelayURL string, resp types.DiscoveryResponse, seenAt time.Time) (RelayState, []RelayState, error) {
101
+ protocolVersion := strings.TrimSpace(resp.ProtocolVersion)
102
+ if protocolVersion != types.ProtocolVersion {
103
+ return RelayState{}, nil, fmt.Errorf("relay protocol version mismatch: relay=%q client=%q", protocolVersion, types.ProtocolVersion)
104
}
24
- if desc.OverlayIPv4 == "" {
25
- return errors.New("descriptor overlay ipv4 is required")
105
+
106
+ selfState, err := newRelayState(resp.Self, seenAt)
107
+ if err != nil {
108
+ return RelayState{}, nil, err
109
+ }
110
+ if err := requireTargetRelayDescriptor(selfState.Descriptor, targetIdentity, targetURL); err != nil {
111
+ return RelayState{}, nil, err
112
+ }
113
+
114
+ relayStates := make([]RelayState, 0, len(resp.Relays))
115
+ seen := map[string]struct{}{selfState.Descriptor.Key(): {}}
116
+ for _, descriptor := range resp.Relays {
117
+ relayState, err := newRelayState(descriptor, seenAt)
118
+ if err != nil {
119
+ continue
120
+ }
121
+ if selfRelay(relayState, selfRelayKey, selfRelayURL) {
122
+ continue
123
+ }
124
+ relayKey := relayState.Descriptor.Key()
125
+ if _, ok := seen[relayKey]; ok {
126
+ continue
127
+ }
128
+ seen[relayKey] = struct{}{}
129
+ relayStates = append(relayStates, relayState)
130
+ }
131
+ return selfState, relayStates, nil
132
+}
133
+
134
+func selfRelay(state RelayState, selfRelayKey string, selfRelayURL string) bool {
135
+ desc := state.Descriptor
136
+ relayKey := desc.Key()
137
+ return (relayKey != "" && selfRelayKey != "" && relayKey == selfRelayKey) ||
138
+ (selfRelayURL != "" && desc.APIHTTPSAddr == selfRelayURL)
139
+}
140
+
141
+func requireTargetRelayDescriptor(desc types.RelayDescriptor, targetIdentity types.Identity, targetURL string) error {
142
+ if strings.TrimSpace(targetIdentity.Name) == "" && strings.TrimSpace(targetIdentity.Address) == "" {
143
+ return errors.New("target relay identity is required")
144
+ }
145
+ targetName := strings.TrimSpace(targetIdentity.Name)
146
+ if targetName != "" {
147
+ normalizedTargetName := utils.NormalizeHostname(targetName)
148
+ if desc.Name != normalizedTargetName {
149
+ return errors.New("descriptor name does not match target relay")
150
+ }
151
+ }
152
+ targetAddress := strings.TrimSpace(targetIdentity.Address)
153
+ if targetAddress != "" {
154
+ normalizedTargetAddress, err := utils.NormalizeEVMAddress(targetAddress)
155
+ if err != nil {
156
+ return err
157
+ }
158
+ if desc.Address != normalizedTargetAddress {
159
+ return errors.New("descriptor address does not match target relay")
160
+ }
161
+ }
162
+ if targetURL != "" {
163
+ normalizedTargetURL, err := utils.NormalizeRelayURL(targetURL)
164
+ if err != nil {
165
+ return err
166
+ }
167
+ if desc.APIHTTPSAddr != normalizedTargetURL {
168
+ return errors.New("descriptor api_https_addr does not match target url")
169
+ }
170
}
171
return nil
172
}
173
+
174
+func (state RelayState) hasDescriptor() bool {
175
+ return !state.LastSeenAt.IsZero() && state.Descriptor.Key() != "" && state.Descriptor.APIHTTPSAddr != ""
176
+}
177
+
178
+type RelayPolicy interface {
179
+ DiscoveryStates(types.Identity, string, string, string, types.DiscoveryResponse, time.Time) (RelayState, []RelayState, error)
180
+ Decide(RelayState) RelayState
181
+ OnDiscovered(RelayState, bool) RelayState
182
+ OnFailure(RelayState, error, int) (RelayState, bool, string)
183
+ OnBanned(RelayState) RelayState
184
+ KeepState(RelayState) bool
185
+ AdvertisedDescriptors([]RelayState) []types.RelayDescriptor
186
+}
187
+
188
+type DefaultRelayPolicy struct{}
189
+
190
+func (DefaultRelayPolicy) Decide(state RelayState) RelayState {
191
+ now := time.Now().UTC()
192
+ status := state.Status
193
+ switch {
194
+ case status == RelayStatusBanned || status == RelayStatusExpired:
195
+ case !state.Descriptor.ExpiresAt.IsZero() && !state.Descriptor.ExpiresAt.After(now):
196
+ status = RelayStatusExpired
197
+ case state.Descriptor.SupportsOverlayPeer &&
198
+ (state.Descriptor.WireGuardPublicKey == "" ||
199
+ state.Descriptor.WireGuardEndpoint == "" ||
200
+ state.Descriptor.OverlayIPv4 == ""):
201
+ status = RelayStatusInvalid
202
+ case status == RelayStatusUnknown:
203
+ status = RelayStatusHinted
204
+ }
205
+
206
+ usable := status == RelayStatusHinted || status == RelayStatusConfirmed
207
+ state.Status = status
208
+ state.Active = (state.Bootstrap && usable) || status == RelayStatusConfirmed
209
+ state.Advertise = status == RelayStatusConfirmed
210
+ state.BootstrapDiscovery = state.Bootstrap && usable
211
+ state.DirectDiscovery = !state.Bootstrap && usable
212
+ state.OverlayDiscovery = !state.Bootstrap && usable && state.Descriptor.SupportsOverlayPeer
213
+ state.OverlayPeer = usable && state.Descriptor.SupportsOverlayPeer
214
+ return state
215
+}
216
+
217
+func (DefaultRelayPolicy) OnDiscovered(state RelayState, advertise bool) RelayState {
218
+ if state.Status == RelayStatusBanned {
219
+ return state
220
+ }
221
+ if advertise {
222
+ state.Status = RelayStatusConfirmed
223
+ state.consecutiveFailures = 0
224
+ return state
225
+ }
226
+ if state.Status != RelayStatusConfirmed {
227
+ state.Status = RelayStatusHinted
228
+ state.consecutiveFailures = 0
229
+ }
230
+ return state
231
+}
232
+
233
+func (DefaultRelayPolicy) OnFailure(state RelayState, err error, recoveryFailures int) (RelayState, bool, string) {
234
+ if state.Status == RelayStatusBanned {
235
+ return state, false, ""
236
+ }
237
+ state.consecutiveFailures++
238
+ if state.Status != RelayStatusExpired && state.consecutiveFailures >= recoveryFailures {
239
+ state.Status = RelayStatusExpired
240
+ return state, true, "recovery"
241
+ }
242
+ var apiErr *types.APIRequestError
243
+ if errors.As(err, &apiErr) &&
244
+ (apiErr.StatusCode == http.StatusForbidden ||
245
+ apiErr.StatusCode == http.StatusNotFound ||
246
+ apiErr.StatusCode == http.StatusGone) {
247
+ state.Status = RelayStatusExpired
248
+ return state, true, "status"
249
+ }
250
+ return state, false, ""
251
+}
252
+
253
+func (DefaultRelayPolicy) OnBanned(state RelayState) RelayState {
254
+ state.Status = RelayStatusBanned
255
+ return state
256
+}
257
+
258
+func (DefaultRelayPolicy) KeepState(state RelayState) bool {
259
+ return state.hasDescriptor() ||
260
+ state.Status == RelayStatusBanned ||
261
+ state.Status == RelayStatusExpired ||
262
+ state.consecutiveFailures > 0
263
+}
264
+
265
+func (p DefaultRelayPolicy) AdvertisedDescriptors(states []RelayState) []types.RelayDescriptor {
266
+ if len(states) == 0 {
267
+ return nil
268
+ }
269
+
270
+ out := make([]types.RelayDescriptor, 0, len(states))
271
+ for _, state := range states {
272
+ state = p.Decide(state)
273
+ if state.Descriptor.APIHTTPSAddr == "" || !state.Advertise {
274
+ continue
275
+ }
276
+ out = append(out, state.Descriptor)
277
+ }
278
+ if len(out) == 0 {
279
+ return nil
280
+ }
281
+ sort.Slice(out, func(i, j int) bool {
282
+ return out[i].APIHTTPSAddr < out[j].APIHTTPSAddr
283
+ })
284
+ return out
285
+}
portal/discovery/refresher.go
+42
-36
@@ -16,12 +16,13 @@ import (
16
)
17
18
const (
19
+ DiscoveryPollInterval = 1 * time.Minute
20
defaultRecoveryFailures = 3
21
)
22
23
type OverlayRuntime interface {
24
DiscoverRelay(context.Context, types.RelayDescriptor) (types.DiscoveryResponse, error)
24
- Sync(map[string]RelayState) error
25
+ Sync([]RelayState) error
26
}
27
28
type Refresher struct {
@@ -67,7 +68,7 @@ func (r *Refresher) Refresh(ctx context.Context) error {
68
if r.overlay == nil {
69
return ctx.Err()
70
}
70
- if err := r.overlay.Sync(r.relaySet.View()); err != nil {
71
+ if err := r.overlay.Sync(r.relaySet.RelayStates()); err != nil {
72
log.Warn().
73
Err(err).
74
Msg("sync wireguard peers")
@@ -77,8 +78,13 @@ func (r *Refresher) Refresh(ctx context.Context) error {
78
}
79
80
func (r *Refresher) refreshHTTPS(ctx context.Context) error {
80
- for _, bootstrap := range r.relaySet.BootstrapDescriptors() {
81
- resp, err := r.discoverHTTPS(ctx, bootstrap)
81
+ states := r.relaySet.RelayStates()
82
+ for _, state := range states {
83
+ if state.Descriptor.APIHTTPSAddr == "" || !state.BootstrapDiscovery {
84
+ continue
85
+ }
86
+ relay := state.Descriptor
87
+ resp, err := r.discoverHTTPS(ctx, relay)
88
if err != nil {
89
if ctx.Err() != nil {
90
return ctx.Err()
@@ -87,7 +93,7 @@ func (r *Refresher) refreshHTTPS(ctx context.Context) error {
93
}
94
95
now := time.Now().UTC()
90
- _, _, err = r.relaySet.ApplyRelayDiscoveryResponse(bootstrap.Identity, bootstrap.APIHTTPSAddr, resp, now)
96
+ _, err = r.relaySet.ApplyRelayDiscoveryResponse(relay.Identity, relay.APIHTTPSAddr, resp, now)
97
if err != nil {
98
continue
99
}
@@ -96,10 +102,15 @@ func (r *Refresher) refreshHTTPS(ctx context.Context) error {
102
return err
103
}
104
99
- for _, relay := range r.relaySet.confirmableDescriptors() {
100
- if r.overlay != nil && relay.SupportsOverlayPeer {
105
+ states = r.relaySet.RelayStates()
106
+ for _, state := range states {
107
+ if state.Descriptor.APIHTTPSAddr == "" || !state.DirectDiscovery {
108
+ continue
109
+ }
110
+ if r.overlay != nil && state.Descriptor.SupportsOverlayPeer {
111
continue
112
}
113
+ relay := state.Descriptor
114
resp, err := r.discoverHTTPS(ctx, relay)
115
if err != nil {
116
if ctx.Err() != nil {
@@ -110,7 +121,7 @@ func (r *Refresher) refreshHTTPS(ctx context.Context) error {
121
}
122
123
now := time.Now().UTC()
113
- _, _, err = r.relaySet.ApplyRelayDiscoveryResponse(relay.Identity, relay.APIHTTPSAddr, resp, now)
124
+ _, err = r.relaySet.ApplyRelayDiscoveryResponse(relay.Identity, relay.APIHTTPSAddr, resp, now)
125
if err != nil {
126
r.logDirectDiscoveryFailure(relay, err, r.directRecoveryFailures)
127
continue
@@ -133,46 +144,41 @@ func (r *Refresher) discoverHTTPS(ctx context.Context, relay types.RelayDescript
144
}
145
146
func (r *Refresher) refreshOverlay(ctx context.Context) error {
136
- for _, relay := range r.relaySet.SyncableDescriptors() {
147
+ for _, state := range r.relaySet.RelayStates() {
148
+ if state.Descriptor.APIHTTPSAddr == "" || !state.OverlayDiscovery {
149
+ continue
150
+ }
151
+
152
+ relay := state.Descriptor
153
var failureErr error
154
139
- if err := RequireOverlayRelayDescriptor(relay); err != nil {
155
+ resp, err := r.overlay.DiscoverRelay(ctx, relay)
156
+ if err != nil {
157
+ if ctx.Err() != nil {
158
+ return ctx.Err()
159
+ }
160
failureErr = err
161
} else {
142
- resp, err := r.overlay.DiscoverRelay(ctx, relay)
143
- if err != nil {
144
- if ctx.Err() != nil {
145
- return ctx.Err()
162
+ now := time.Now().UTC()
163
+ relaySetChanged, err := r.relaySet.ApplyRelayDiscoveryResponse(relay.Identity, relay.APIHTTPSAddr, resp, now)
164
+ if relaySetChanged {
165
+ if syncErr := r.overlay.Sync(r.relaySet.RelayStates()); syncErr != nil {
166
+ log.Warn().
167
+ Err(syncErr).
168
+ Str("relay", relay.APIHTTPSAddr).
169
+ Msg("sync wireguard peers")
170
}
171
+ }
172
+ if err != nil {
173
failureErr = err
174
} else {
149
- now := time.Now().UTC()
150
- relaySetChanged, warnErr, err := r.relaySet.ApplyOverlayRelayDiscoveryResponse(relay.Identity, relay.APIHTTPSAddr, resp, now)
151
- if relaySetChanged {
152
- view := r.relaySet.View()
153
- if syncErr := r.overlay.Sync(view); syncErr != nil {
154
- if warnErr == nil {
155
- warnErr = syncErr
156
- }
157
- }
158
- }
159
- if err != nil {
160
- failureErr = err
161
- } else {
162
- if warnErr != nil {
163
- log.Warn().
164
- Err(warnErr).
165
- Str("relay", relay.APIHTTPSAddr).
166
- Msg("overlay relay discovery completed with warnings")
167
- }
168
- continue
169
- }
175
+ continue
176
}
177
}
178
179
expired, expireReason, consecutiveFailures := r.relaySet.RecordDiscoveryFailure(relay.Identity, relay.APIHTTPSAddr, failureErr, r.overlayRecoveryFailures)
180
if expired {
175
- if syncErr := r.overlay.Sync(r.relaySet.View()); syncErr != nil && failureErr == nil {
181
+ if syncErr := r.overlay.Sync(r.relaySet.RelayStates()); syncErr != nil && failureErr == nil {
182
failureErr = syncErr
183
}
184
}
portal/discovery/relayset.go
+151
-592
@@ -2,8 +2,6 @@ package discovery
2
3
import (
4
"errors"
5
- "fmt"
6
- "net/http"
5
"reflect"
6
"slices"
7
"sort"
@@ -11,76 +9,36 @@ import (
9
"sync"
10
"time"
11
14
- "github.com/rs/zerolog/log"
15
-
12
"github.com/gosuda/portal-tunnel/v2/types"
13
"github.com/gosuda/portal-tunnel/v2/utils"
14
)
15
20
-type relayStatus uint8
21
-
22
-const (
23
- relayStatusHinted relayStatus = iota
24
- relayStatusConfirmed
25
- relayStatusExpired
26
-)
27
-
28
-// RelayState is the single relay shape shared by discovery storage and overlay sync.
29
-// Package-private fields are internal-only local state.
30
-type RelayState struct {
31
- Descriptor types.RelayDescriptor
32
- FirstSeenAt time.Time
33
- LastSeenAt time.Time
34
- Banned bool
35
- Expired bool
36
- status relayStatus
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
-}
56
-
16
// RelaySet owns the shared relay discovery view: configured bootstrap relay URLs,
17
// the latest validated descriptor seen for each relay, and local runtime state
18
// such as ban/reachability/failure tracking.
19
type RelaySet struct {
20
mu sync.RWMutex
21
knownRelayURLs []string
63
- relayKeysByURL map[string]string
22
relays map[string]RelayState
65
- localByURL map[string]RelayState
23
+ policy RelayPolicy
24
selfRelayKey string
25
selfRelayURL string
26
}
27
70
-type relayDescriptorProjection struct {
71
- state RelayState
72
- relayURL string
73
- bootstrap bool
74
-}
75
-
28
func NewRelaySet(identity types.Identity, relayURL string, bootstrapRelayURLs []string) (*RelaySet, error) {
77
- set := &RelaySet{
78
- relayKeysByURL: make(map[string]string),
79
- relays: make(map[string]RelayState),
80
- localByURL: make(map[string]RelayState),
29
+ if relayURL != "" {
30
+ normalized, err := utils.NormalizeRelayURL(relayURL)
31
+ if err != nil {
32
+ return nil, err
33
+ }
34
+ relayURL = normalized
35
}
82
- if err := set.SetSelfRelay(identity, relayURL); err != nil {
83
- return nil, err
36
+
37
+ set := &RelaySet{
38
+ relays: make(map[string]RelayState),
39
+ policy: DefaultRelayPolicy{},
40
+ selfRelayKey: identity.Key(),
41
+ selfRelayURL: relayURL,
42
}
43
if err := set.SetBootstrapRelayURLs(bootstrapRelayURLs); err != nil {
44
return nil, err
@@ -88,180 +46,81 @@ func NewRelaySet(identity types.Identity, relayURL string, bootstrapRelayURLs []
46
return set, nil
47
}
48
91
-func relayExpiredAt(state RelayState, now time.Time) bool {
92
- if state.status == relayStatusExpired {
93
- return true
94
- }
95
- if state.Descriptor.ExpiresAt.IsZero() {
96
- return false
97
- }
98
- if now.IsZero() {
99
- now = time.Now().UTC()
49
+func (s *RelaySet) SetRelayPolicy(policy RelayPolicy) {
50
+ if policy == nil {
51
+ policy = DefaultRelayPolicy{}
52
}
101
- return !state.Descriptor.ExpiresAt.After(now)
102
-}
103
-
104
-func (s *RelaySet) isSelfRelayURLLocked(relayURL string) bool {
105
- relayURL = strings.TrimSpace(relayURL)
106
- return relayURL != "" && s.selfRelayURL != "" && relayURL == s.selfRelayURL
107
-}
108
-
109
-func (s *RelaySet) isSelfRelayDescriptorLocked(desc types.RelayDescriptor) bool {
110
- if relayKey := desc.Key(); relayKey != "" && s.selfRelayKey != "" && relayKey == s.selfRelayKey {
111
- return true
112
- }
113
- return s.isSelfRelayURLLocked(desc.APIHTTPSAddr)
114
-}
115
-
116
-func (s *RelaySet) storeLocalStateLocked(relayURL string, state RelayState) {
117
- relayURL = strings.TrimSpace(relayURL)
118
- if relayURL == "" {
119
- return
120
- }
121
- if state.isDefaultLocalState() {
122
- delete(s.localByURL, relayURL)
123
- return
124
- }
125
- s.localByURL[relayURL] = state
53
+ s.mu.Lock()
54
+ defer s.mu.Unlock()
55
+ s.policy = policy
56
}
57
128
-func (s *RelaySet) SetSelfRelay(identity types.Identity, relayURL string) error {
129
- relayURL = strings.TrimSpace(relayURL)
130
- if relayURL != "" {
131
- normalized, err := utils.NormalizeRelayURL(relayURL)
132
- if err != nil {
133
- return err
134
- }
135
- relayURL = normalized
58
+func (s *RelaySet) SetBootstrapRelayURLs(inputs []string) error {
59
+ normalized, err := utils.NormalizeRelayURLs(inputs...)
60
+ if err != nil {
61
+ return err
62
}
63
64
s.mu.Lock()
65
defer s.mu.Unlock()
140
- s.selfRelayKey = identity.Key()
141
- s.selfRelayURL = relayURL
142
- if s.selfRelayURL != "" {
143
- filtered := s.knownRelayURLs[:0]
144
- for _, knownRelayURL := range s.knownRelayURLs {
145
- if s.isSelfRelayURLLocked(knownRelayURL) {
146
- continue
147
- }
148
- filtered = append(filtered, knownRelayURL)
149
- }
150
- s.knownRelayURLs = filtered
151
- delete(s.localByURL, s.selfRelayURL)
152
- delete(s.relayKeysByURL, s.selfRelayURL)
153
- }
154
- if s.selfRelayKey != "" {
155
- if record, ok := s.relays[s.selfRelayKey]; ok {
156
- delete(s.localByURL, record.Descriptor.APIHTTPSAddr)
157
- delete(s.relayKeysByURL, record.Descriptor.APIHTTPSAddr)
158
- }
159
- delete(s.relays, s.selfRelayKey)
160
- }
161
- return nil
162
-}
66
164
-func (s *RelaySet) bootstrapRelayURLsLocked() []string {
165
- if len(s.knownRelayURLs) == 0 {
166
- return nil
167
- }
168
-
169
- out := make([]string, 0, len(s.knownRelayURLs))
170
- for _, relayURL := range s.knownRelayURLs {
171
- if s.isSelfRelayURLLocked(relayURL) || s.localByURL[relayURL].Banned {
172
- continue
173
- }
174
- out = append(out, relayURL)
175
- }
176
- if len(out) == 0 {
177
- return nil
178
- }
179
- return out
180
-}
67
+ filtered := utils.RemoveRelayURL(normalized, s.selfRelayURL)
68
182
-func (s *RelaySet) descriptorProjectionsLocked() []relayDescriptorProjection {
183
- if len(s.relays) == 0 {
184
- return nil
69
+ keep := make(map[string]struct{}, len(filtered))
70
+ for _, relayURL := range filtered {
71
+ keep[relayURL] = struct{}{}
72
}
73
187
- out := make([]relayDescriptorProjection, 0, len(s.relays))
188
- for _, record := range s.relays {
189
- if s.isSelfRelayDescriptorLocked(record.Descriptor) {
74
+ for _, relayURL := range s.knownRelayURLs {
75
+ if _, ok := keep[relayURL]; ok {
76
continue
77
}
192
- relayURL := strings.TrimSpace(record.Descriptor.APIHTTPSAddr)
193
- if relayURL == "" {
78
+ state := s.relays[relayURL]
79
+ if state.hasDescriptor() {
80
continue
81
}
196
- bootstrap := false
197
- if !s.isSelfRelayURLLocked(relayURL) {
198
- for _, candidate := range s.knownRelayURLs {
199
- if candidate == relayURL {
200
- bootstrap = true
201
- break
202
- }
203
- }
82
+ if !s.policy.KeepState(state) {
83
+ delete(s.relays, relayURL)
84
}
205
- local := s.localByURL[relayURL]
206
- record.Banned = local.Banned
207
- record.status = local.status
208
- record.consecutiveFailures = local.consecutiveFailures
209
- out = append(out, relayDescriptorProjection{
210
- state: record,
211
- relayURL: relayURL,
212
- bootstrap: bootstrap,
213
- })
214
- }
215
- if len(out) == 0 {
216
- return nil
85
}
218
- sort.Slice(out, func(i, j int) bool {
219
- return out[i].relayURL < out[j].relayURL
220
- })
221
- return out
222
-}
86
224
-func (s *RelaySet) ActiveRelayURLs() []string {
225
- s.mu.RLock()
226
- defer s.mu.RUnlock()
87
+ s.knownRelayURLs = append([]string(nil), filtered...)
88
+ return nil
89
+}
90
228
- return selectRelayURLsLocked(s.relayCandidatesLocked(time.Now().UTC()), SelectionOptions{})
91
+func (s *RelaySet) relayURLForKeyLocked(relayKey string) string {
92
+ if relayKey == "" {
93
+ return ""
94
+ }
95
+ for relayURL, state := range s.relays {
96
+ if state.hasDescriptor() && state.Descriptor.Key() == relayKey {
97
+ return relayURL
98
+ }
99
+ }
100
+ return ""
101
}
102
231
-func (s *RelaySet) RelayCandidates() []RelayCandidate {
103
+func (s *RelaySet) RelayStates() []RelayState {
104
s.mu.RLock()
105
defer s.mu.RUnlock()
106
235
- return s.relayCandidatesLocked(time.Now().UTC())
107
+ states := s.relayStatesLocked()
108
+ for i := range states {
109
+ states[i] = s.policy.Decide(states[i])
110
+ }
111
+ return states
112
}
113
238
-func (s *RelaySet) SelectRelayURLs(opts SelectionOptions) []string {
114
+func (s *RelaySet) ActiveRelays() []RelayState {
115
s.mu.RLock()
116
defer s.mu.RUnlock()
117
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
118
+ states := s.relayStatesLocked()
119
+ out := make([]RelayState, 0, len(states))
120
+ for _, state := range states {
121
+ state = s.policy.Decide(state)
122
+ if state.Active && state.Descriptor.APIHTTPSAddr != "" {
123
+ out = append(out, state)
124
}
125
}
126
if len(out) == 0 {
@@ -270,449 +129,167 @@ func selectRelayURLsLocked(candidates []RelayCandidate, opts SelectionOptions) [
129
return out
130
}
131
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
-}
132
+func (s *RelaySet) relayStatesLocked() []RelayState {
133
+ out := make([]RelayState, 0, len(s.knownRelayURLs)+len(s.relays))
134
+ seen := make(map[string]struct{}, len(s.knownRelayURLs)+len(s.relays))
135
281
-func (s *RelaySet) relayCandidatesLocked(now time.Time) []RelayCandidate {
282
- bootstrapRelayURLs := s.bootstrapRelayURLsLocked()
283
- projections := s.descriptorProjectionsLocked()
284
-
285
- out := make([]RelayCandidate, 0, len(bootstrapRelayURLs)+len(projections))
286
- seen := make(map[string]struct{}, len(bootstrapRelayURLs)+len(projections))
287
- for _, relayURL := range bootstrapRelayURLs {
136
+ for _, relayURL := range s.knownRelayURLs {
137
+ if relayURL == "" {
138
+ continue
139
+ }
140
if _, ok := seen[relayURL]; ok {
141
continue
142
}
143
seen[relayURL] = struct{}{}
144
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
- },
145
+ state, ok := s.relays[relayURL]
146
+ if !ok {
147
+ state = newRelayHintState(relayURL)
148
}
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
- }
149
+ if selfRelay(state, s.selfRelayKey, s.selfRelayURL) {
150
+ continue
151
}
311
- out = append(out, candidate)
152
+ state.Bootstrap = true
153
+ out = append(out, state)
154
}
313
- for _, projection := range projections {
314
- if projection.state.Banned || projection.state.status != relayStatusConfirmed || relayExpiredAt(projection.state, now) {
155
+
156
+ descriptorStates := make([]RelayState, 0, len(s.relays))
157
+ for relayURL, record := range s.relays {
158
+ if !record.hasDescriptor() {
159
continue
160
}
317
- if _, ok := seen[projection.relayURL]; ok {
161
+ relayURL = strings.TrimSpace(record.Descriptor.APIHTTPSAddr)
162
+ if relayURL == "" {
163
continue
164
}
320
- seen[projection.relayURL] = struct{}{}
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
330
- }
331
- return out
332
-}
333
-
334
-func (s *RelaySet) BootstrapDescriptors() []types.RelayDescriptor {
335
- s.mu.RLock()
336
- defer s.mu.RUnlock()
337
-
338
- bootstrapRelayURLs := s.bootstrapRelayURLsLocked()
339
- if len(bootstrapRelayURLs) == 0 {
340
- return nil
341
- }
342
-
343
- out := make([]types.RelayDescriptor, 0, len(bootstrapRelayURLs))
344
- for _, relayURL := range bootstrapRelayURLs {
345
- if relayKey, ok := s.relayKeysByURL[relayURL]; ok {
346
- if record, ok := s.relays[relayKey]; ok && record.Descriptor.APIHTTPSAddr != "" {
347
- out = append(out, record.Descriptor)
348
- continue
349
- }
350
- }
351
- out = append(out, types.RelayDescriptor{
352
- Identity: types.Identity{
353
- Name: utils.PortalRootHost(relayURL),
354
- },
355
- RelayID: relayURL,
356
- APIHTTPSAddr: relayURL,
357
- Version: 1,
358
- })
359
- }
360
- if len(out) == 0 {
361
- return nil
362
- }
363
- return out
364
-}
365
-
366
-func (s *RelaySet) BanRelayURL(relayURL string) {
367
- s.mu.Lock()
368
- defer s.mu.Unlock()
369
- relayURL = strings.TrimSpace(relayURL)
370
- if relayURL == "" {
371
- return
372
- }
373
-
374
- state := s.localByURL[relayURL]
375
- state.Banned = true
376
- s.storeLocalStateLocked(relayURL, state)
377
-}
378
-
379
-func (s *RelaySet) AdvertisedDescriptors() []types.RelayDescriptor {
380
- s.mu.RLock()
381
- defer s.mu.RUnlock()
382
-
383
- now := time.Now().UTC()
384
- projections := s.descriptorProjectionsLocked()
385
- if len(projections) == 0 {
386
- return nil
387
- }
388
-
389
- out := make([]types.RelayDescriptor, 0, len(projections))
390
- for _, projection := range projections {
391
- if projection.state.Banned || projection.state.status != relayStatusConfirmed || relayExpiredAt(projection.state, now) {
165
+ if selfRelay(record, s.selfRelayKey, s.selfRelayURL) {
166
continue
167
}
394
- out = append(out, projection.state.Descriptor)
395
- }
396
- if len(out) == 0 {
397
- return nil
398
- }
399
- return out
400
-}
401
-
402
-func (s *RelaySet) confirmableDescriptors() []types.RelayDescriptor {
403
- s.mu.RLock()
404
- defer s.mu.RUnlock()
405
-
406
- now := time.Now().UTC()
407
- projections := s.descriptorProjectionsLocked()
408
- if len(projections) == 0 {
409
- return nil
410
- }
411
-
412
- out := make([]types.RelayDescriptor, 0, len(projections))
413
- for _, projection := range projections {
414
- if projection.bootstrap || projection.state.Banned || relayExpiredAt(projection.state, now) {
168
+ if _, ok := seen[relayURL]; ok {
169
continue
170
}
417
- out = append(out, projection.state.Descriptor)
418
- }
419
- if len(out) == 0 {
420
- return nil
421
- }
422
- return out
423
-}
171
425
-func (s *RelaySet) SyncableDescriptors() []types.RelayDescriptor {
426
- s.mu.RLock()
427
- defer s.mu.RUnlock()
428
-
429
- now := time.Now().UTC()
430
- projections := s.descriptorProjectionsLocked()
431
- if len(projections) == 0 {
432
- return nil
172
+ record.Bootstrap = slices.Contains(s.knownRelayURLs, relayURL)
173
+ descriptorStates = append(descriptorStates, record)
174
}
175
+ sort.Slice(descriptorStates, func(i, j int) bool {
176
+ return descriptorStates[i].Descriptor.APIHTTPSAddr < descriptorStates[j].Descriptor.APIHTTPSAddr
177
+ })
178
+ out = append(out, descriptorStates...)
179
435
- out := make([]types.RelayDescriptor, 0, len(projections))
436
- for _, projection := range projections {
437
- if projection.bootstrap || projection.state.Banned || relayExpiredAt(projection.state, now) || !projection.state.Descriptor.SupportsOverlayPeer {
438
- continue
439
- }
440
- out = append(out, projection.state.Descriptor)
441
- }
180
if len(out) == 0 {
181
return nil
182
}
183
return out
184
}
185
448
-func (s *RelaySet) View() map[string]RelayState {
186
+func (s *RelaySet) AdvertisedDescriptors() []types.RelayDescriptor {
187
s.mu.RLock()
188
defer s.mu.RUnlock()
189
452
- now := time.Now().UTC()
453
- projections := s.descriptorProjectionsLocked()
454
- if len(projections) == 0 {
455
- return nil
456
- }
457
-
458
- view := make(map[string]RelayState, len(projections))
459
- for _, projection := range projections {
460
- relayKey := projection.state.Descriptor.Key()
461
- if relayKey == "" {
462
- continue
463
- }
464
- expired := relayExpiredAt(projection.state, now)
465
- view[relayKey] = RelayState{
466
- Descriptor: projection.state.Descriptor,
467
- FirstSeenAt: projection.state.FirstSeenAt,
468
- LastSeenAt: projection.state.LastSeenAt,
469
- Banned: projection.state.Banned,
470
- Expired: expired,
471
- }
472
- }
473
- if len(view) == 0 {
474
- return nil
475
- }
476
- return view
190
+ return s.policy.AdvertisedDescriptors(s.relayStatesLocked())
191
}
192
479
-func (s *RelaySet) SetBootstrapRelayURLs(inputs []string) error {
480
- normalized, err := utils.NormalizeRelayURLs(inputs...)
481
- if err != nil {
482
- return err
483
- }
484
-
193
+func (s *RelaySet) BanRelayURL(relayURL string) {
194
s.mu.Lock()
195
defer s.mu.Unlock()
487
-
488
- filtered := utils.RemoveRelayURL(normalized, s.selfRelayURL)
489
-
490
- keep := make(map[string]struct{}, len(filtered))
491
- for _, relayURL := range filtered {
492
- keep[relayURL] = struct{}{}
196
+ relayURL = strings.TrimSpace(relayURL)
197
+ if relayURL == "" {
198
+ return
199
}
200
495
- for _, relayURL := range s.knownRelayURLs {
496
- if _, ok := keep[relayURL]; ok {
497
- continue
498
- }
499
- if _, ok := s.relayKeysByURL[relayURL]; ok {
500
- continue
501
- }
502
- if state := s.localByURL[relayURL]; state.isDefaultLocalState() {
503
- delete(s.localByURL, relayURL)
504
- }
201
+ state := s.relays[relayURL]
202
+ if state.Descriptor.APIHTTPSAddr == "" {
203
+ state = newRelayHintState(relayURL)
204
}
506
-
507
- s.knownRelayURLs = append([]string(nil), filtered...)
508
- return nil
205
+ state = s.policy.OnBanned(state)
206
+ s.relays[relayURL] = state
207
}
208
511
-func (s *RelaySet) storeDescriptorLocked(desc types.RelayDescriptor, now time.Time) (string, bool, bool, error) {
209
+func (s *RelaySet) storeDescriptorLocked(state RelayState, policy RelayPolicy) (bool, bool, error) {
210
+ desc := state.Descriptor
211
relayKey := desc.Key()
212
if relayKey == "" {
514
- return "", false, false, errors.New("descriptor identity is required")
213
+ return false, false, errors.New("descriptor identity is required")
214
}
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() {
520
- now = time.Now().UTC()
215
+ if existing := s.relays[desc.APIHTTPSAddr]; existing.hasDescriptor() && existing.Descriptor.Key() != relayKey {
216
+ return false, false, errors.New("descriptor identity does not match known relay url")
217
}
218
523
- record, ok := s.relays[relayKey]
524
- added := !ok
525
- if !ok {
526
- record.FirstSeenAt = now
527
- }
528
- previousURL := record.Descriptor.APIHTTPSAddr
219
+ previousURL := s.relayURLForKeyLocked(relayKey)
220
+ record := s.relays[desc.APIHTTPSAddr]
221
previousDescriptor := record.Descriptor
530
- record.Descriptor = desc
531
- record.LastSeenAt = now
532
- s.relays[relayKey] = record
533
- s.relayKeysByURL[desc.APIHTTPSAddr] = relayKey
222
if previousURL != "" && previousURL != desc.APIHTTPSAddr {
535
- delete(s.relayKeysByURL, previousURL)
536
- state := s.localByURL[previousURL]
537
- bootstrap := false
538
- if !s.isSelfRelayURLLocked(previousURL) {
539
- if slices.Contains(s.knownRelayURLs, previousURL) {
540
- bootstrap = true
541
- }
223
+ previous := s.relays[previousURL]
224
+ previousDescriptor = previous.Descriptor
225
+ if !policy.KeepState(record) {
226
+ record.Status = previous.Status
227
+ record.consecutiveFailures = previous.consecutiveFailures
228
}
543
- if !bootstrap && state.isDefaultLocalState() {
544
- delete(s.localByURL, previousURL)
229
+ delete(s.relays, previousURL)
230
+ }
231
+
232
+ added := previousURL == ""
233
+ if !policy.KeepState(record) {
234
+ record = state
235
+ } else {
236
+ if record.FirstSeenAt.IsZero() {
237
+ record.FirstSeenAt = state.FirstSeenAt
238
}
239
+ record.Descriptor = desc
240
+ record.LastSeenAt = state.LastSeenAt
241
}
242
+ s.relays[desc.APIHTTPSAddr] = record
243
244
changed := added || !reflect.DeepEqual(previousDescriptor, desc)
549
- return relayKey, added, changed, nil
245
+ return added, changed, nil
246
}
247
552
-func (s *RelaySet) applyDiscoveryResponse(targetIdentity types.Identity, targetURL string, resp types.DiscoveryResponse, now time.Time, requireOverlay bool) (relaySetChanged bool, warnErr error, err error) {
248
+func (s *RelaySet) ApplyRelayDiscoveryResponse(targetIdentity types.Identity, targetURL string, resp types.DiscoveryResponse, now time.Time) (relaySetChanged bool, err error) {
249
if now.IsZero() {
250
now = time.Now().UTC()
251
} else {
252
now = now.UTC()
253
}
254
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
- }
255
+ s.mu.Lock()
256
+ defer s.mu.Unlock()
257
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
- }
258
+ policy := s.policy
259
596
- selfDescriptor, err := normalize(resp.Self)
260
+ selfState, relayStates, err := policy.DiscoveryStates(targetIdentity, targetURL, s.selfRelayKey, s.selfRelayURL, resp, now)
261
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
- }
262
+ return false, err
263
}
264
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
- }
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
-
676
- _, added, descriptorChanged, err := s.storeDescriptorLocked(desc, now)
265
+ apply := func(state RelayState, advertise bool) error {
266
+ desc := state.Descriptor
267
+ added, descriptorChanged, err := s.storeDescriptorLocked(state, policy)
268
if err != nil {
269
return err
270
}
271
681
- localState := s.localByURL[desc.APIHTTPSAddr]
682
- previousState := localState
683
- if advertise {
684
- localState.status = relayStatusConfirmed
685
- localState.consecutiveFailures = 0
686
- } else if localState.status != relayStatusConfirmed {
687
- localState.status = relayStatusHinted
688
- localState.consecutiveFailures = 0
689
- }
690
- s.storeLocalStateLocked(desc.APIHTTPSAddr, localState)
272
+ storedState := s.relays[desc.APIHTTPSAddr]
273
+ previousState := storedState
274
+ storedState = policy.OnDiscovered(storedState, advertise)
275
+ s.relays[desc.APIHTTPSAddr] = storedState
276
692
- changed := added || descriptorChanged || !reflect.DeepEqual(previousState, localState)
277
+ changed := added || descriptorChanged || !reflect.DeepEqual(previousState, s.relays[desc.APIHTTPSAddr])
278
if changed {
279
relaySetChanged = true
280
}
281
return nil
282
}
283
699
- if err := apply(selfDescriptor, true); err != nil {
700
- return false, warnErr, err
284
+ if err := apply(selfState, true); err != nil {
285
+ return false, err
286
}
702
- for _, relayDescriptor := range relayDescriptors {
703
- if err := apply(relayDescriptor, false); err != nil {
704
- return false, warnErr, err
287
+ for _, relayState := range relayStates {
288
+ if err := apply(relayState, false); err != nil {
289
+ return false, err
290
}
291
}
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) {
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) {
715
- return s.applyDiscoveryResponse(targetIdentity, targetURL, resp, now, true)
292
+ return relaySetChanged, nil
293
}
294
295
func (s *RelaySet) RecordDiscoveryFailure(identity types.Identity, relayURL string, err error, recoveryFailures int) (expired bool, expireReason string, consecutiveFailures int) {
@@ -725,36 +302,18 @@ func (s *RelaySet) RecordDiscoveryFailure(identity types.Identity, relayURL stri
302
s.mu.Lock()
303
defer s.mu.Unlock()
304
728
- record, ok := s.relays[relayKey]
729
- if !ok {
730
- return false, "", 0
731
- }
732
- if relayURL == "" || s.relayKeysByURL[relayURL] != relayKey {
733
- relayURL = record.Descriptor.APIHTTPSAddr
305
+ if relayURL == "" || s.relays[relayURL].Descriptor.Key() != relayKey {
306
+ relayURL = s.relayURLForKeyLocked(relayKey)
307
}
308
if relayURL == "" {
309
return false, "", 0
310
}
311
739
- localState := s.localByURL[relayURL]
740
- localState.consecutiveFailures++
741
- s.storeLocalStateLocked(relayURL, localState)
742
- if localState.status != relayStatusExpired && localState.consecutiveFailures >= recoveryFailures {
743
- state := s.localByURL[record.Descriptor.APIHTTPSAddr]
744
- state.status = relayStatusExpired
745
- s.storeLocalStateLocked(record.Descriptor.APIHTTPSAddr, state)
746
- return true, "recovery", localState.consecutiveFailures
747
- }
748
-
749
- var apiErr *types.APIRequestError
750
- if errors.As(err, &apiErr) &&
751
- (apiErr.StatusCode == http.StatusForbidden ||
752
- apiErr.StatusCode == http.StatusNotFound ||
753
- apiErr.StatusCode == http.StatusGone) {
754
- state := s.localByURL[record.Descriptor.APIHTTPSAddr]
755
- state.status = relayStatusExpired
756
- s.storeLocalStateLocked(record.Descriptor.APIHTTPSAddr, state)
757
- return true, "status", localState.consecutiveFailures
312
+ state, ok := s.relays[relayURL]
313
+ if !ok {
314
+ return false, "", 0
315
}
759
- return false, "", localState.consecutiveFailures
316
+ state, expired, expireReason = s.policy.OnFailure(state, err, recoveryFailures)
317
+ s.relays[relayURL] = state
318
+ return expired, expireReason, state.consecutiveFailures
319
}
portal/overlay/overlay.go
+8
-10
@@ -209,24 +209,22 @@ func (o *Overlay) DiscoverRelay(ctx context.Context, relay types.RelayDescriptor
209
return resp, nil
210
}
211
212
-func (o *Overlay) Sync(view map[string]discovery.RelayState) error {
212
+func (o *Overlay) Sync(relays []discovery.RelayState) error {
213
if o == nil || o.stack == nil {
214
return nil
215
}
216
- return o.stack.ApplyPeers(peersForView(o.cfg.PublicKey, view))
216
+ return o.stack.ApplyPeers(peersForRelays(o.cfg.PublicKey, relays))
217
}
218
219
-func peersForView(publicKey string, view map[string]discovery.RelayState) []desiredPeer {
220
- peers := make([]desiredPeer, 0, len(view))
221
- for _, relay := range view {
222
- if relay.Expired || relay.Banned {
219
+func peersForRelays(publicKey string, relays []discovery.RelayState) []desiredPeer {
220
+ peers := make([]desiredPeer, 0, len(relays))
221
+ for _, relay := range relays {
222
+ if !relay.OverlayPeer {
223
continue
224
}
225
+
226
desc := relay.Descriptor
226
- if desc.WireGuardPublicKey == publicKey || !desc.SupportsOverlayPeer {
227
- continue
228
- }
229
- if desc.WireGuardPublicKey == "" || desc.WireGuardEndpoint == "" || desc.OverlayIPv4 == "" {
227
+ if desc.WireGuardPublicKey == publicKey {
228
continue
229
}
230
portal/server.go
+1
-1
@@ -614,7 +614,7 @@ func (s *Server) startOverlay() (*overlay.Overlay, error) {
614
return nil, fmt.Errorf("start wireguard overlay: %w", err)
615
}
616
617
- if err := overlay.Sync(s.relaySet.View()); err != nil {
617
+ if err := overlay.Sync(s.relaySet.RelayStates()); err != nil {
618
_ = overlay.Shutdown(context.Background())
619
return nil, fmt.Errorf("sync wireguard peers: %w", err)
620
}
portal/server_test.go
+20
-8
@@ -51,11 +51,8 @@ func mustRelayDescriptor(t *testing.T, relayURL string) types.RelayDescriptor {
51
52
func applyRelay(t *testing.T, set *discovery.RelaySet, identity types.Identity, targetURL string, resp types.DiscoveryResponse, now time.Time) error {
53
t.Helper()
54
- _, warnErr, err := set.ApplyRelayDiscoveryResponse(identity, targetURL, resp, now)
55
- if err != nil {
56
- return err
57
- }
58
- return warnErr
54
+ _, err := set.ApplyRelayDiscoveryResponse(identity, targetURL, resp, now)
55
+ return err
56
}
57
58
func tempIdentityPath(t *testing.T) string {
@@ -511,7 +508,12 @@ func TestServerSetBootstrapRelayURLsAllowsLoopbackButSkipsSelfRelay(t *testing.T
508
t.Fatalf("SetBootstrapRelayURLs() error = %v", err)
509
}
510
advertisedDescriptors := server.relaySet.AdvertisedDescriptors()
514
- knownURLs := append([]string(nil), server.relaySet.ActiveRelayURLs()...)
511
+ knownURLs := make([]string, 0)
512
+ for _, state := range server.relaySet.RelayStates() {
513
+ if state.Active && state.Descriptor.APIHTTPSAddr != "" {
514
+ knownURLs = append(knownURLs, state.Descriptor.APIHTTPSAddr)
515
+ }
516
+ }
517
sort.Strings(knownURLs)
518
if !reflect.DeepEqual(knownURLs, []string{
519
"https://bootstrap.example.com",
@@ -564,7 +566,12 @@ func TestServerDiscoverySkipsSelfRelayHint(t *testing.T) {
566
t.Fatalf("ApplyRelayDiscoveryResponse() error = %v", err)
567
}
568
567
- knownURLs := append([]string(nil), server.relaySet.ActiveRelayURLs()...)
569
+ knownURLs := make([]string, 0)
570
+ for _, state := range server.relaySet.RelayStates() {
571
+ if state.Active && state.Descriptor.APIHTTPSAddr != "" {
572
+ knownURLs = append(knownURLs, state.Descriptor.APIHTTPSAddr)
573
+ }
574
+ }
575
sort.Strings(knownURLs)
576
if !reflect.DeepEqual(knownURLs, []string{"https://bootstrap.example.com"}) {
577
t.Fatalf("ActiveRelayURLs() = %v, want self hint excluded", knownURLs)
@@ -609,7 +616,12 @@ func TestServerRecordVerifiedDiscoveryPeerRequiresDirectConfirmation(t *testing.
616
if err != nil {
617
t.Fatalf("applyRelayDiscoveryResponse() hinted error = %v", err)
618
}
612
- knownURLs := append([]string(nil), server.relaySet.ActiveRelayURLs()...)
619
+ knownURLs := make([]string, 0)
620
+ for _, state := range server.relaySet.RelayStates() {
621
+ if state.Active && state.Descriptor.APIHTTPSAddr != "" {
622
+ knownURLs = append(knownURLs, state.Descriptor.APIHTTPSAddr)
623
+ }
624
+ }
625
sort.Strings(knownURLs)
626
if !reflect.DeepEqual(knownURLs, []string{"https://bootstrap.example.com"}) {
627
t.Fatalf("ActiveRelayURLs() = %v, want [%q]", knownURLs, "https://bootstrap.example.com")
sdk/expose.go
+11
-5
@@ -159,7 +159,11 @@ func (e *Exposure) runDiscoveryLoop(ctx context.Context) {
159
}
160
161
func (e *Exposure) ActiveRelayURLs() []string {
162
- return e.relaySet.ActiveRelayURLs()
162
+ var out []string
163
+ for _, state := range e.relaySet.ActiveRelays() {
164
+ out = append(out, state.Descriptor.APIHTTPSAddr)
165
+ }
166
+ return out
167
}
168
169
func (e *Exposure) Addr() net.Addr {
@@ -254,11 +258,10 @@ func (e *Exposure) WaitDatagramReady(ctx context.Context) ([]string, error) {
258
259
func (e *Exposure) RunHTTP(ctx context.Context, handler http.Handler, localAddr string) error {
260
var relayListener net.Listener
257
- activeRelayURLs := e.relaySet.ActiveRelayURLs()
261
e.listenerMu.RLock()
262
activeListeners := make([]*Listener, 0, len(e.relayListeners))
260
- for _, relayURL := range activeRelayURLs {
261
- listener, ok := e.relayListeners[relayURL]
263
+ for _, state := range e.relaySet.ActiveRelays() {
264
+ listener, ok := e.relayListeners[state.Descriptor.APIHTTPSAddr]
265
if !ok {
266
continue
267
}
@@ -365,12 +368,15 @@ func (e *Exposure) Close() error {
368
}
369
370
func (e *Exposure) reconcileRelayListeners(failOnError bool) error {
368
- activeRelayURLs := e.relaySet.ActiveRelayURLs()
371
e.listenerMu.Lock()
372
currentRelayURLs := make([]string, 0, len(e.relayListeners))
373
for relayURL := range e.relayListeners {
374
currentRelayURLs = append(currentRelayURLs, relayURL)
375
}
376
+ activeRelayURLs := make([]string, 0, len(e.relayListeners))
377
+ for _, state := range e.relaySet.ActiveRelays() {
378
+ activeRelayURLs = append(activeRelayURLs, state.Descriptor.APIHTTPSAddr)
379
+ }
380
missingRelayURLs := utils.FilterRelayURLs(activeRelayURLs, currentRelayURLs)
381
staleRelayURLs := utils.FilterRelayURLs(currentRelayURLs, activeRelayURLs)
382
staleListeners := make([]*Listener, 0, len(staleRelayURLs))
sdk/expose_test.go
+5
-8
@@ -43,11 +43,8 @@ func mustRelayDescriptor(t *testing.T, relayName, relayURL string) types.RelayDe
43
44
func applyRelayDiscovery(t *testing.T, set *discovery.RelaySet, identity types.Identity, targetURL string, resp types.DiscoveryResponse, now time.Time) error {
45
t.Helper()
46
- _, warnErr, err := set.ApplyRelayDiscoveryResponse(identity, targetURL, resp, now)
47
- if err != nil {
48
- return err
49
- }
50
- return warnErr
46
+ _, err := set.ApplyRelayDiscoveryResponse(identity, targetURL, resp, now)
47
+ return err
48
}
49
50
func TestExposureBanRelayURLMovesRelay(t *testing.T) {
@@ -83,7 +80,7 @@ func TestExposureBanRelayURLMovesRelay(t *testing.T) {
80
t.Fatalf("ActiveRelayURLs() = %v, want [%q]", got, relayB)
81
}
82
86
- knownRelayURLs := exposure.relaySet.ActiveRelayURLs()
83
+ knownRelayURLs := exposure.ActiveRelayURLs()
84
exposure.listenerMu.RLock()
85
_, listenerExists := exposure.relayListeners[relayA]
86
exposure.listenerMu.RUnlock()
@@ -119,7 +116,7 @@ func TestExposureSetRelayURLsSkipsBannedRelay(t *testing.T) {
116
if got := exposure.ActiveRelayURLs(); len(got) != 1 || got[0] != relayA {
117
t.Fatalf("ActiveRelayURLs() = %v, want [%q]", got, relayA)
118
}
122
- knownRelayURLs := exposure.relaySet.ActiveRelayURLs()
119
+ knownRelayURLs := exposure.ActiveRelayURLs()
120
if len(knownRelayURLs) != 1 || knownRelayURLs[0] != relayA {
121
t.Fatalf("knownRelayURLs = %v, want [%q]", knownRelayURLs, relayA)
122
}
@@ -169,7 +166,7 @@ func TestExposureSetRelayURLsRemovesStaleListener(t *testing.T) {
166
t.Fatal("stale relay listener was not closed")
167
}
168
172
- knownRelayURLs := exposure.relaySet.ActiveRelayURLs()
169
+ knownRelayURLs := exposure.ActiveRelayURLs()
170
exposure.listenerMu.RLock()
171
_, relayAExists := exposure.relayListeners[relayA]
172
_, relayBExists := exposure.relayListeners[relayB]
sdk/mitm_test.go
+8
-3
@@ -228,8 +228,8 @@ func TestMITMProbeDetectionBansListener(t *testing.T) {
228
Reason: types.MITMProbeReasonExporterMismatch,
229
}, nil)
230
231
- for _, activeRelayURL := range listener.relaySet.ActiveRelayURLs() {
232
- if activeRelayURL == relayURL.String() {
231
+ for _, state := range listener.relaySet.RelayStates() {
232
+ if state.Active && state.Descriptor.APIHTTPSAddr == relayURL.String() {
233
t.Fatal("relay still active after mitm detection")
234
}
235
}
@@ -260,7 +260,12 @@ func TestMITMProbeDetectionWarnsWithoutBanningListener(t *testing.T) {
260
Reason: types.MITMProbeReasonExporterMismatch,
261
}, nil)
262
263
- activeRelayURLs := listener.relaySet.ActiveRelayURLs()
263
+ activeRelayURLs := make([]string, 0)
264
+ for _, state := range listener.relaySet.RelayStates() {
265
+ if state.Active && state.Descriptor.APIHTTPSAddr != "" {
266
+ activeRelayURLs = append(activeRelayURLs, state.Descriptor.APIHTTPSAddr)
267
+ }
268
+ }
269
if len(activeRelayURLs) != 1 || activeRelayURLs[0] != relayURL.String() {
270
t.Fatalf("ActiveRelayURLs() = %v, want [%q]", activeRelayURLs, relayURL.String())
271
}