Refactor relay discovery and state management
- Updated the relay discovery logic to improve handling of expired and banned relays. - Enhanced the `RelaySet` structure to include methods for confirming and unconfirming relays. - Modified the `Refresher` to skip banned relays during the refresh process. - Adjusted the relay state management to ensure that stale hints are preserved within a retention window. - Removed redundant checks and improved the clarity of the relay selection logic.
Kim committed
Apr 13, 2026 at 21:08 UTC
b7a889876f9a04e84d3a2082ecce193d7ba9b1c0
11 files changed
+747
-157
portal/api_server.go
+1
-1
@@ -176,7 +176,7 @@ func (s *Server) handleRelayDiscovery(w http.ResponseWriter, r *http.Request) {
176
OwnerAddress: s.identity.Address,
177
Version: 1,
178
IssuedAt: now,
179
- ExpiresAt: now.Add(2 * discovery.DiscoveryPollInterval),
179
+ ExpiresAt: now.Add(discovery.DiscoveryDescriptorTTL),
180
APIHTTPSAddr: s.cfg.PortalURL,
181
Discovery: s.cfg.DiscoveryEnabled,
182
IngressTLSAddr: ingressAddr,
portal/discovery/policy.go
+49
-43
@@ -11,10 +11,11 @@ import (
11
)
12
13
type RelayPolicy interface {
14
- SelectActive([]RelayState) []RelayState
14
+ SelectAggregate([]RelayState) []RelayState
15
+ SelectConfirmed([]RelayState) []RelayState
16
SelectPriority([]RelayState, ClientState) []string
17
OnConfirmed(RelayState) RelayState
17
- OnHinted(RelayState) RelayState
18
+ OnUnconfirmed(RelayState) RelayState
19
OnFailure(RelayState, error, int) (RelayState, bool, string)
20
OnBanned(RelayState) RelayState
21
}
@@ -23,14 +24,13 @@ type DefaultRelayPolicy struct{}
24
25
const highDiscoveryRTTThreshold = 1 * time.Second
26
26
-func (p DefaultRelayPolicy) selectStates(states []RelayState, keep func(RelayState) bool) []RelayState {
27
- now := time.Now().UTC()
27
+func (p DefaultRelayPolicy) SelectAggregate(states []RelayState) []RelayState {
28
out := make([]RelayState, 0, len(states))
29
for _, state := range states {
30
- if !state.discoverable(now) {
30
+ if state.Banned {
31
continue
32
}
33
- if !keep(state) {
33
+ if !state.Bootstrap && !state.hasDescriptor() {
34
continue
35
}
36
out = append(out, state)
@@ -41,14 +41,23 @@ func (p DefaultRelayPolicy) selectStates(states []RelayState, keep func(RelaySta
41
return out
42
}
43
44
-func (p DefaultRelayPolicy) SelectActive(states []RelayState) []RelayState {
45
- return p.selectStates(states, func(state RelayState) bool {
46
- return state.Bootstrap || state.Confirmed
47
- })
44
+func (p DefaultRelayPolicy) SelectConfirmed(states []RelayState) []RelayState {
45
+ selected := p.SelectAggregate(states)
46
+ out := make([]RelayState, 0, len(selected))
47
+ for _, state := range selected {
48
+ if !state.Confirmed {
49
+ continue
50
+ }
51
+ out = append(out, state)
52
+ }
53
+ if len(out) == 0 {
54
+ return nil
55
+ }
56
+ return out
57
}
58
59
func (p DefaultRelayPolicy) SelectPriority(states []RelayState, clientState ClientState) []string {
51
- selected := p.SelectActive(states)
60
+ selected := p.SelectAggregate(states)
61
if len(selected) == 0 {
62
return nil
63
}
@@ -74,46 +83,46 @@ func (p DefaultRelayPolicy) SelectPriority(states []RelayState, clientState Clie
83
}
84
85
currentAuto := make([]string, 0, len(autoPool))
77
- remainingAuto := make([]string, 0, len(autoPool))
86
+ confirmedAuto := make([]string, 0, len(autoPool))
87
highRTTAuto := make([]string, 0, len(autoPool))
79
- penalizedAuto := make([]string, 0, len(autoPool))
88
+ remainingAuto := make([]string, 0, len(autoPool))
89
for _, state := range autoPool {
90
relayURL := state.Descriptor.APIHTTPSAddr
82
- statePenalized := state.consecutiveFailures > 0 && !state.Reachable
91
switch {
84
- case slices.Contains(clientState.ActiveRelayURLs, relayURL) && !statePenalized:
92
+ case slices.Contains(clientState.ActiveRelayURLs, relayURL):
93
currentAuto = append(currentAuto, relayURL)
86
- case statePenalized:
87
- penalizedAuto = append(penalizedAuto, relayURL)
88
- case !state.DiscoveryRTTAt.IsZero() && state.DiscoveryRTT > highDiscoveryRTTThreshold:
94
+ case state.Confirmed && !state.DiscoveryRTTAt.IsZero() && state.DiscoveryRTT > highDiscoveryRTTThreshold:
95
highRTTAuto = append(highRTTAuto, relayURL)
96
+ case state.Confirmed:
97
+ confirmedAuto = append(confirmedAuto, relayURL)
98
default:
99
remainingAuto = append(remainingAuto, relayURL)
100
}
101
}
102
103
+ if len(confirmedAuto) > 1 {
104
+ rng := rand.New(rand.NewSource(time.Now().UnixNano()))
105
+ rng.Shuffle(len(confirmedAuto), func(i, j int) {
106
+ confirmedAuto[i], confirmedAuto[j] = confirmedAuto[j], confirmedAuto[i]
107
+ })
108
+ }
109
rng := rand.New(rand.NewSource(time.Now().UnixNano()))
110
if len(remainingAuto) > 1 {
111
rng.Shuffle(len(remainingAuto), func(i, j int) {
112
remainingAuto[i], remainingAuto[j] = remainingAuto[j], remainingAuto[i]
113
})
114
}
101
- if len(penalizedAuto) > 1 {
102
- rng.Shuffle(len(penalizedAuto), func(i, j int) {
103
- penalizedAuto[i], penalizedAuto[j] = penalizedAuto[j], penalizedAuto[i]
104
- })
105
- }
115
if len(highRTTAuto) > 1 {
116
rng.Shuffle(len(highRTTAuto), func(i, j int) {
117
highRTTAuto[i], highRTTAuto[j] = highRTTAuto[j], highRTTAuto[i]
118
})
119
}
120
112
- autoURLs := make([]string, 0, len(currentAuto)+len(remainingAuto)+len(highRTTAuto)+len(penalizedAuto))
121
+ autoURLs := make([]string, 0, len(currentAuto)+len(confirmedAuto)+len(highRTTAuto)+len(remainingAuto))
122
autoURLs = append(autoURLs, currentAuto...)
114
- autoURLs = append(autoURLs, remainingAuto...)
123
+ autoURLs = append(autoURLs, confirmedAuto...)
124
autoURLs = append(autoURLs, highRTTAuto...)
116
- autoURLs = append(autoURLs, penalizedAuto...)
125
+ autoURLs = append(autoURLs, remainingAuto...)
126
if clientState.MaxActiveRelays > 0 && len(autoURLs) > clientState.MaxActiveRelays {
127
autoURLs = autoURLs[:clientState.MaxActiveRelays]
128
}
@@ -131,20 +140,12 @@ func (DefaultRelayPolicy) OnConfirmed(state RelayState) RelayState {
140
if state.Banned {
141
return state
142
}
134
- state.Reachable = true
143
state.Confirmed = true
136
- state.consecutiveFailures = 0
144
return state
145
}
146
140
-func (DefaultRelayPolicy) OnHinted(state RelayState) RelayState {
141
- if state.Banned {
142
- return state
143
- }
144
- state.Reachable = true
145
- if !state.Confirmed {
146
- state.consecutiveFailures = 0
147
- }
147
+func (DefaultRelayPolicy) OnUnconfirmed(state RelayState) RelayState {
148
+ state.Confirmed = false
149
return state
150
}
151
@@ -153,12 +154,16 @@ func (DefaultRelayPolicy) OnFailure(state RelayState, err error, recoveryFailure
154
return state, false, ""
155
}
156
state.consecutiveFailures++
156
- expire := func(reason string) (RelayState, bool, string) {
157
- if !state.Reachable {
158
- return state, false, ""
157
+ backOff := func(reason string) (RelayState, bool, string) {
158
+ backoff := defaultDirectRecoveryBackoff
159
+ for extra := state.consecutiveFailures - recoveryFailures; extra > 0; extra-- {
160
+ if backoff >= maxDirectRecoveryBackoff/2 {
161
+ backoff = maxDirectRecoveryBackoff
162
+ break
163
+ }
164
+ backoff *= 2
165
}
160
- state.Reachable = false
161
- state.Confirmed = false
166
+ state.nextDirectRefreshAt = time.Now().UTC().Add(backoff)
167
return state, true, reason
168
}
169
var apiErr *types.APIRequestError
@@ -166,15 +171,16 @@ func (DefaultRelayPolicy) OnFailure(state RelayState, err error, recoveryFailure
171
(apiErr.StatusCode == http.StatusForbidden ||
172
apiErr.StatusCode == http.StatusNotFound ||
173
apiErr.StatusCode == http.StatusGone) {
169
- return expire("status")
174
+ return backOff("status")
175
}
176
if state.consecutiveFailures >= recoveryFailures {
172
- return expire("recovery")
177
+ return backOff("recovery")
178
}
179
return state, false, ""
180
}
181
182
func (DefaultRelayPolicy) OnBanned(state RelayState) RelayState {
183
state.Banned = true
184
+ state.Confirmed = false
185
return state
186
}
portal/discovery/policy_test.go
+180
-19
@@ -1,6 +1,7 @@
1
package discovery
2
3
import (
4
+ "errors"
5
"testing"
6
"time"
7
@@ -21,6 +22,7 @@ func mustPolicyRelayDescriptor(t *testing.T, relayName, relayURL string) types.R
22
IssuedAt: now,
23
ExpiresAt: now.Add(time.Hour),
24
APIHTTPSAddr: relayURL,
25
+ Discovery: true,
26
})
27
if err != nil {
28
t.Fatalf("NormalizeDescriptor() error = %v", err)
@@ -46,7 +48,6 @@ func confirmedPolicyRelayState(t *testing.T, relayName, relayURL string) RelaySt
48
49
return RelayState{
50
Descriptor: mustPolicyRelayDescriptor(t, relayName, relayURL),
49
- Reachable: true,
51
Confirmed: true,
52
LastSeenAt: time.Now().UTC(),
53
}
@@ -84,6 +85,93 @@ func TestSelectPriorityKeepsExplicitRelaysOutsideAutoLimit(t *testing.T) {
85
}
86
}
87
88
+func TestSelectAggregateKeepsBootstrapRelayWhenDescriptorExpired(t *testing.T) {
89
+ policy := DefaultRelayPolicy{}
90
+ relayURL := "https://relay-bootstrap.example"
91
+
92
+ state := bootstrapPolicyRelayState(relayURL)
93
+ state.LastSeenAt = time.Now().UTC().Add(-time.Minute)
94
+ state.Descriptor.ExpiresAt = time.Now().UTC().Add(-time.Second)
95
+
96
+ selected := policy.SelectAggregate([]RelayState{state})
97
+
98
+ if len(selected) != 1 {
99
+ t.Fatalf("len(selected) = %d, want 1", len(selected))
100
+ }
101
+ if got := selected[0].Descriptor.APIHTTPSAddr; got != relayURL {
102
+ t.Fatalf("selected[0] = %q, want bootstrap relay %q", got, relayURL)
103
+ }
104
+}
105
+
106
+func TestSelectAggregateKeepsCollectedRelayEvenWhenNotAdvertisable(t *testing.T) {
107
+ policy := DefaultRelayPolicy{}
108
+ state := RelayState{
109
+ Descriptor: mustPolicyRelayDescriptor(t, "relay-a", "https://relay-a.example"),
110
+ LastSeenAt: time.Now().UTC().Add(-DiscoveryHintRetentionTTL).Add(-time.Hour),
111
+ }
112
+ state.Descriptor.Discovery = false
113
+ state.Descriptor.ExpiresAt = time.Now().UTC().Add(-time.Second)
114
+
115
+ selected := policy.SelectAggregate([]RelayState{state})
116
+
117
+ if len(selected) != 1 {
118
+ t.Fatalf("len(selected) = %d, want 1", len(selected))
119
+ }
120
+}
121
+
122
+func TestSelectAggregateIncludesHintedRelayWithoutConfirmation(t *testing.T) {
123
+ policy := DefaultRelayPolicy{}
124
+ state := RelayState{
125
+ Descriptor: mustPolicyRelayDescriptor(t, "relay-hinted", "https://relay-hinted.example"),
126
+ LastSeenAt: time.Now().UTC(),
127
+ }
128
+
129
+ selected := policy.SelectAggregate([]RelayState{state})
130
+
131
+ if len(selected) != 1 {
132
+ t.Fatalf("len(selected) = %d, want 1", len(selected))
133
+ }
134
+ if got := selected[0].Descriptor.APIHTTPSAddr; got != state.Descriptor.APIHTTPSAddr {
135
+ t.Fatalf("selected[0] = %q, want %q", got, state.Descriptor.APIHTTPSAddr)
136
+ }
137
+}
138
+
139
+func TestSelectAggregateSkipsBannedBootstrapRelay(t *testing.T) {
140
+ policy := DefaultRelayPolicy{}
141
+ state := bootstrapPolicyRelayState("https://relay-banned.example")
142
+ state.Banned = true
143
+
144
+ selected := policy.SelectAggregate([]RelayState{state})
145
+
146
+ if len(selected) != 0 {
147
+ t.Fatalf("len(selected) = %d, want 0", len(selected))
148
+ }
149
+}
150
+
151
+func TestSelectConfirmedKeepsOnlyConfirmedAggregateRelays(t *testing.T) {
152
+ policy := DefaultRelayPolicy{}
153
+ confirmed := confirmedPolicyRelayState(t, "relay-confirmed", "https://relay-confirmed.example")
154
+ hinted := RelayState{
155
+ Descriptor: mustPolicyRelayDescriptor(t, "relay-hinted", "https://relay-hinted.example"),
156
+ LastSeenAt: time.Now().UTC(),
157
+ }
158
+ bannedConfirmed := confirmedPolicyRelayState(t, "relay-banned", "https://relay-banned.example")
159
+ bannedConfirmed.Banned = true
160
+
161
+ selected := policy.SelectConfirmed([]RelayState{
162
+ hinted,
163
+ confirmed,
164
+ bannedConfirmed,
165
+ })
166
+
167
+ if len(selected) != 1 {
168
+ t.Fatalf("len(selected) = %d, want 1", len(selected))
169
+ }
170
+ if got := selected[0].Descriptor.APIHTTPSAddr; got != confirmed.Descriptor.APIHTTPSAddr {
171
+ t.Fatalf("selected[0] = %q, want confirmed relay %q", got, confirmed.Descriptor.APIHTTPSAddr)
172
+ }
173
+}
174
+
175
func TestSelectPriorityColdStartSelectsEligibleRelay(t *testing.T) {
176
policy := DefaultRelayPolicy{}
177
relayA := "https://relay-a.example"
@@ -104,7 +192,30 @@ func TestSelectPriorityColdStartSelectsEligibleRelay(t *testing.T) {
192
}
193
}
194
107
-func TestSelectPriorityKeepsCurrentHealthyRelayOverNewConfirmedRelay(t *testing.T) {
195
+func TestSelectPriorityPrefersConfirmedRelayOverHintedRelay(t *testing.T) {
196
+ policy := DefaultRelayPolicy{}
197
+ confirmedRelay := confirmedPolicyRelayState(t, "relay-confirmed", "https://relay-confirmed.example")
198
+ hintedRelay := RelayState{
199
+ Descriptor: mustPolicyRelayDescriptor(t, "relay-hinted", "https://relay-hinted.example"),
200
+ LastSeenAt: time.Now().UTC(),
201
+ }
202
+
203
+ selected := policy.SelectPriority([]RelayState{
204
+ hintedRelay,
205
+ confirmedRelay,
206
+ }, ClientState{
207
+ MaxActiveRelays: 1,
208
+ })
209
+
210
+ if len(selected) != 1 {
211
+ t.Fatalf("len(selected) = %d, want 1", len(selected))
212
+ }
213
+ if got := selected[0]; got != confirmedRelay.Descriptor.APIHTTPSAddr {
214
+ t.Fatalf("selected[0] = %q, want confirmed relay %q", got, confirmedRelay.Descriptor.APIHTTPSAddr)
215
+ }
216
+}
217
+
218
+func TestSelectPriorityKeepsCurrentRelayOverNewConfirmedRelay(t *testing.T) {
219
policy := DefaultRelayPolicy{}
220
currentRelay := "https://relay-current.example"
221
newRelay := "https://relay-new.example"
@@ -121,7 +232,7 @@ func TestSelectPriorityKeepsCurrentHealthyRelayOverNewConfirmedRelay(t *testing.
232
t.Fatalf("len(selected) = %d, want 1", len(selected))
233
}
234
if got := selected[0]; got != currentRelay {
124
- t.Fatalf("selected[0] = %q, want current healthy relay %q kept", got, currentRelay)
235
+ t.Fatalf("selected[0] = %q, want current relay %q kept", got, currentRelay)
236
}
237
}
238
@@ -145,25 +256,75 @@ func TestSelectPriorityPushesHighRTTRelayBehindNormalRelay(t *testing.T) {
256
}
257
}
258
148
-func TestSelectPriorityReplacesCurrentDeadRelay(t *testing.T) {
259
+func TestOnConfirmedMarksRelayConfirmed(t *testing.T) {
260
policy := DefaultRelayPolicy{}
150
- currentRelay := bootstrapPolicyRelayState("https://relay-current.example")
151
- currentRelay.Reachable = false
152
- currentRelay.consecutiveFailures = 1
261
+ nextDirectRefreshAt := time.Now().UTC().Add(time.Minute)
262
+ state := RelayState{
263
+ Descriptor: mustPolicyRelayDescriptor(t, "relay-a", "https://relay-a.example"),
264
+ LastSeenAt: time.Now().UTC(),
265
+ consecutiveFailures: defaultRecoveryFailures,
266
+ nextDirectRefreshAt: nextDirectRefreshAt,
267
+ }
268
154
- replacementRelay := confirmedPolicyRelayState(t, "relay-new", "https://relay-new.example")
155
- selected := policy.SelectPriority([]RelayState{
156
- currentRelay,
157
- replacementRelay,
158
- }, ClientState{
159
- ActiveRelayURLs: []string{currentRelay.Descriptor.APIHTTPSAddr},
160
- MaxActiveRelays: 1,
161
- })
269
+ state = policy.OnConfirmed(state)
270
163
- if len(selected) != 1 {
164
- t.Fatalf("len(selected) = %d, want 1", len(selected))
271
+ if !state.Confirmed {
272
+ t.Fatal("relay should become confirmed")
273
+ }
274
+ if state.consecutiveFailures != defaultRecoveryFailures {
275
+ t.Fatalf("consecutiveFailures = %d, want %d", state.consecutiveFailures, defaultRecoveryFailures)
276
+ }
277
+ if !state.nextDirectRefreshAt.Equal(nextDirectRefreshAt) {
278
+ t.Fatalf("nextDirectRefreshAt = %v, want %v", state.nextDirectRefreshAt, nextDirectRefreshAt)
279
}
166
- if got := selected[0]; got != replacementRelay.Descriptor.APIHTTPSAddr {
167
- t.Fatalf("selected[0] = %q, want replacement relay %q", got, replacementRelay.Descriptor.APIHTTPSAddr)
280
+}
281
+
282
+func TestOnUnconfirmedClearsRelayConfirmation(t *testing.T) {
283
+ policy := DefaultRelayPolicy{}
284
+ state := confirmedPolicyRelayState(t, "relay-a", "https://relay-a.example")
285
+
286
+ state = policy.OnUnconfirmed(state)
287
+
288
+ if state.Confirmed {
289
+ t.Fatal("relay should become unconfirmed")
290
+ }
291
+}
292
+
293
+func TestOnFailureSchedulesDirectRecoveryRetry(t *testing.T) {
294
+ policy := DefaultRelayPolicy{}
295
+ state := confirmedPolicyRelayState(t, "relay-a", "https://relay-a.example")
296
+ startedAt := time.Now().UTC()
297
+
298
+ var backedOff bool
299
+ var reason string
300
+ for range defaultRecoveryFailures {
301
+ state, backedOff, reason = policy.OnFailure(state, errors.New("boom"), defaultRecoveryFailures)
302
+ }
303
+
304
+ if !backedOff {
305
+ t.Fatal("expected relay to back off after recovery failure budget")
306
+ }
307
+ if reason != "recovery" {
308
+ t.Fatalf("backoff reason = %q, want recovery", reason)
309
+ }
310
+ if !state.nextDirectRefreshAt.After(startedAt) {
311
+ t.Fatalf("nextDirectRefreshAt = %v, want a future retry time", state.nextDirectRefreshAt)
312
+ }
313
+}
314
+
315
+func TestOnFailureSchedulesRetryForHintedRelay(t *testing.T) {
316
+ policy := DefaultRelayPolicy{}
317
+ state := RelayState{
318
+ Descriptor: mustPolicyRelayDescriptor(t, "relay-hinted", "https://relay-hinted.example"),
319
+ LastSeenAt: time.Now().UTC(),
320
+ }
321
+ startedAt := time.Now().UTC()
322
+
323
+ for range defaultRecoveryFailures {
324
+ state, _, _ = policy.OnFailure(state, errors.New("boom"), defaultRecoveryFailures)
325
+ }
326
+
327
+ if !state.nextDirectRefreshAt.After(startedAt) {
328
+ t.Fatalf("nextDirectRefreshAt = %v, want a future retry time", state.nextDirectRefreshAt)
329
}
330
}
portal/discovery/refresher.go
+19
-7
@@ -19,7 +19,7 @@ import (
19
const (
20
defaultRequestTimeout = 15 * time.Second
21
DiscoveryPollInterval = 30 * time.Second
22
- defaultRecoveryFailures = 3
22
+ defaultRecoveryFailures = 5
23
)
24
25
type OverlayRuntime interface {
@@ -98,7 +98,19 @@ func (r *Refresher) refreshHTTPS(ctx context.Context, extraSourceHosts []string)
98
99
now := time.Now().UTC()
100
for _, state := range states {
101
- if !state.discoverable(now) || (state.hasDescriptor() && !state.Descriptor.Discovery) {
101
+ if state.Banned {
102
+ continue
103
+ }
104
+ if !state.hasDescriptor() {
105
+ if !state.Bootstrap {
106
+ continue
107
+ }
108
+ } else if !state.Bootstrap {
109
+ if !state.nextDirectRefreshAt.IsZero() && state.nextDirectRefreshAt.After(now) {
110
+ continue
111
+ }
112
+ }
113
+ if state.hasDescriptor() && !state.Descriptor.Discovery {
114
continue
115
}
116
@@ -226,18 +238,18 @@ func (r *Refresher) refreshOverlay(ctx context.Context) error {
238
}
239
240
func (r *Refresher) logDiscoveryFailure(targetRelayURL, sourceURL string, recoveryFailures int, err error) {
229
- expired, expireReason, consecutiveFailures := r.relaySet.RecordRelayFailure(targetRelayURL, err, recoveryFailures)
230
- if !expired {
241
+ backedOff, backoffReason, consecutiveFailures := r.relaySet.RecordRelayFailure(targetRelayURL, err, recoveryFailures)
242
+ if !backedOff {
243
return
244
}
245
246
event := log.Warn().
247
Err(err).
248
Str("relay", sourceURL).
237
- Bool("expired", true).
238
- Str("reason", expireReason)
249
+ Bool("backed_off", true).
250
+ Str("reason", backoffReason)
251
if consecutiveFailures > 0 {
252
event = event.Int("consecutive_failures", consecutiveFailures)
253
}
242
- event.Msg("discovery source expired")
254
+ event.Msg("discovery source retry delayed")
255
}
portal/discovery/refresher_test.go
new
+154
@@ -0,0 +1,154 @@
1
+package discovery
2
+
3
+import (
4
+ "context"
5
+ "encoding/pem"
6
+ "net/http"
7
+ "net/http/httptest"
8
+ "sync/atomic"
9
+ "testing"
10
+ "time"
11
+
12
+ "github.com/gosuda/portal-tunnel/v2/types"
13
+ "github.com/gosuda/portal-tunnel/v2/utils"
14
+)
15
+
16
+func newDiscoveryTestServer(t *testing.T, handler http.HandlerFunc) (*httptest.Server, []byte) {
17
+ t.Helper()
18
+
19
+ server := httptest.NewTLSServer(handler)
20
+ rootCAPEM := pem.EncodeToMemory(&pem.Block{Type: "CERTIFICATE", Bytes: server.Certificate().Raw})
21
+ if rootCAPEM == nil {
22
+ server.Close()
23
+ t.Fatal("failed to encode test certificate")
24
+ }
25
+ return server, rootCAPEM
26
+}
27
+
28
+func TestRefresherRefreshesExpiredBootstrapRelay(t *testing.T) {
29
+ now := time.Now().UTC()
30
+ var relayURL string
31
+ server, rootCAPEM := newDiscoveryTestServer(t, http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) {
32
+ utils.WriteAPIData(w, http.StatusOK, types.DiscoveryResponse{
33
+ ProtocolVersion: types.DiscoveryVersion,
34
+ GeneratedAt: now,
35
+ Relays: []types.RelayDescriptor{
36
+ mustPolicyRelayDescriptor(t, "relay-bootstrap", relayURL),
37
+ },
38
+ })
39
+ }))
40
+ defer server.Close()
41
+ relayURL = server.URL
42
+
43
+ set, err := NewRelaySet([]string{relayURL})
44
+ if err != nil {
45
+ t.Fatalf("NewRelaySet() error = %v", err)
46
+ }
47
+
48
+ set.mu.Lock()
49
+ state := set.relays[relayURL]
50
+ state.Descriptor = mustPolicyRelayDescriptor(t, "relay-bootstrap", relayURL)
51
+ state.Descriptor.ExpiresAt = now.Add(-time.Second)
52
+ state.LastSeenAt = now.Add(-time.Minute)
53
+ set.relays[relayURL] = state
54
+ set.mu.Unlock()
55
+
56
+ refresher, err := NewRefresher(set, rootCAPEM, nil, "")
57
+ if err != nil {
58
+ t.Fatalf("NewRefresher() error = %v", err)
59
+ }
60
+ if err := refresher.Refresh(context.Background()); err != nil {
61
+ t.Fatalf("Refresh() error = %v", err)
62
+ }
63
+
64
+ set.mu.RLock()
65
+ refreshed := set.relays[relayURL]
66
+ set.mu.RUnlock()
67
+ if refreshed.Confirmed {
68
+ t.Fatal("direct discovery refresh should not mark relay locally confirmed")
69
+ }
70
+ if !refreshed.Descriptor.ExpiresAt.After(now) {
71
+ t.Fatalf("refreshed descriptor expiry = %v, want a fresh descriptor", refreshed.Descriptor.ExpiresAt)
72
+ }
73
+}
74
+
75
+func TestRefresherRefreshesExpiredCollectedRelay(t *testing.T) {
76
+ now := time.Now().UTC()
77
+ var relayURL string
78
+ server, rootCAPEM := newDiscoveryTestServer(t, http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) {
79
+ utils.WriteAPIData(w, http.StatusOK, types.DiscoveryResponse{
80
+ ProtocolVersion: types.DiscoveryVersion,
81
+ GeneratedAt: now,
82
+ Relays: []types.RelayDescriptor{
83
+ mustPolicyRelayDescriptor(t, "relay-known", relayURL),
84
+ },
85
+ })
86
+ }))
87
+ defer server.Close()
88
+ relayURL = server.URL
89
+
90
+ set, err := NewRelaySet(nil)
91
+ if err != nil {
92
+ t.Fatalf("NewRelaySet() error = %v", err)
93
+ }
94
+
95
+ set.mu.Lock()
96
+ state := confirmedPolicyRelayState(t, "relay-known", relayURL)
97
+ state.Descriptor.ExpiresAt = now.Add(-time.Second)
98
+ state.LastSeenAt = now.Add(-time.Minute)
99
+ set.relays[relayURL] = state
100
+ set.mu.Unlock()
101
+
102
+ refresher, err := NewRefresher(set, rootCAPEM, nil, "")
103
+ if err != nil {
104
+ t.Fatalf("NewRefresher() error = %v", err)
105
+ }
106
+ if err := refresher.Refresh(context.Background()); err != nil {
107
+ t.Fatalf("Refresh() error = %v", err)
108
+ }
109
+
110
+ set.mu.RLock()
111
+ refreshed := set.relays[relayURL]
112
+ set.mu.RUnlock()
113
+ if !refreshed.Descriptor.ExpiresAt.After(now) {
114
+ t.Fatalf("refreshed descriptor expiry = %v, want a fresh descriptor", refreshed.Descriptor.ExpiresAt)
115
+ }
116
+}
117
+
118
+func TestRefresherSkipsDirectRetryUntilBackoffExpires(t *testing.T) {
119
+ now := time.Now().UTC()
120
+ var requests atomic.Int32
121
+ server, rootCAPEM := newDiscoveryTestServer(t, http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) {
122
+ requests.Add(1)
123
+ utils.WriteAPIData(w, http.StatusOK, types.DiscoveryResponse{
124
+ ProtocolVersion: types.DiscoveryVersion,
125
+ GeneratedAt: now,
126
+ Relays: []types.RelayDescriptor{
127
+ mustPolicyRelayDescriptor(t, "relay-known", r.Host),
128
+ },
129
+ })
130
+ }))
131
+ defer server.Close()
132
+
133
+ set, err := NewRelaySet(nil)
134
+ if err != nil {
135
+ t.Fatalf("NewRelaySet() error = %v", err)
136
+ }
137
+
138
+ set.mu.Lock()
139
+ state := confirmedPolicyRelayState(t, "relay-known", server.URL)
140
+ state.nextDirectRefreshAt = now.Add(time.Minute)
141
+ set.relays[server.URL] = state
142
+ set.mu.Unlock()
143
+
144
+ refresher, err := NewRefresher(set, rootCAPEM, nil, "")
145
+ if err != nil {
146
+ t.Fatalf("NewRefresher() error = %v", err)
147
+ }
148
+ if err := refresher.Refresh(context.Background()); err != nil {
149
+ t.Fatalf("Refresh() error = %v", err)
150
+ }
151
+ if got := requests.Load(); got != 0 {
152
+ t.Fatalf("direct refresh requests = %d, want 0 before backoff expires", got)
153
+ }
154
+}
portal/discovery/relayset.go
+69
-25
@@ -15,7 +15,7 @@ import (
15
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 and observed discovery RTT.
18
+// such as ban/failure tracking and observed discovery RTT.
19
type RelaySet struct {
20
mu sync.RWMutex
21
relays map[string]RelayState
@@ -74,11 +74,18 @@ func (s *RelaySet) SetBootstrapRelayURLs(inputs []string) error {
74
return nil
75
}
76
77
-func (s *RelaySet) ActiveRelays() []RelayState {
77
+func (s *RelaySet) AggregateRelays() []RelayState {
78
s.mu.RLock()
79
defer s.mu.RUnlock()
80
81
- return s.policy.SelectActive(s.relayStatesLocked())
81
+ return s.policy.SelectAggregate(s.relayStatesLocked())
82
+}
83
+
84
+func (s *RelaySet) ConfirmedRelays() []RelayState {
85
+ s.mu.RLock()
86
+ defer s.mu.RUnlock()
87
+
88
+ return s.policy.SelectConfirmed(s.relayStatesLocked())
89
}
90
91
func (s *RelaySet) PriorityRelays(clientState ClientState) []string {
@@ -96,7 +103,7 @@ func (s *RelaySet) OverlayPeerStates() []RelayState {
103
now := time.Now().UTC()
104
out := make([]RelayState, 0, len(states))
105
for _, state := range states {
99
- if !state.discoverable(now) || !state.Descriptor.SupportsOverlayPeer {
106
+ if state.Banned || !state.hasDescriptor() || !state.Descriptor.ExpiresAt.After(now) || !state.Descriptor.SupportsOverlayPeer {
107
continue
108
}
109
if state.Descriptor.WireGuardPublicKey == "" ||
@@ -120,10 +127,21 @@ func (s *RelaySet) Descriptors() []types.RelayDescriptor {
127
now := time.Now().UTC()
128
out := make([]types.RelayDescriptor, 0, len(states))
129
for _, state := range states {
123
- if !state.hasDescriptor() || !state.Descriptor.ExpiresAt.After(now) || !state.Descriptor.Discovery {
130
+ if state.Banned || !state.hasDescriptor() || !state.Descriptor.Discovery {
131
continue
132
}
126
- out = append(out, state.Descriptor)
133
+ desc := state.Descriptor
134
+ if !desc.ExpiresAt.After(now) {
135
+ if state.LastSeenAt.IsZero() || !state.LastSeenAt.After(now.Add(-DiscoveryHintRetentionTTL)) {
136
+ continue
137
+ }
138
+
139
+ // Keep stale relay hints flowing through discovery so the mesh converges
140
+ // on a large shared relay set. Local listener confirmation and direct
141
+ // refresh retry state are tracked separately.
142
+ desc.ExpiresAt = now.Add(DiscoveryDescriptorTTL)
143
+ }
144
+ out = append(out, desc)
145
}
146
if len(out) == 0 {
147
return nil
@@ -191,32 +209,53 @@ func (s *RelaySet) BanRelayURL(relayURL string) {
209
s.relays[relayURL] = state
210
}
211
212
+func (s *RelaySet) ConfirmRelayURL(relayURL string) {
213
+ s.mu.Lock()
214
+ defer s.mu.Unlock()
215
+
216
+ state, ok := s.relays[relayURL]
217
+ if !ok {
218
+ state = newRelayStateFromURL(relayURL)
219
+ }
220
+ state = s.policy.OnConfirmed(state)
221
+ s.relays[relayURL] = state
222
+}
223
+
224
+func (s *RelaySet) UnconfirmRelayURL(relayURL string) {
225
+ s.mu.Lock()
226
+ defer s.mu.Unlock()
227
+
228
+ state, ok := s.relays[relayURL]
229
+ if !ok {
230
+ return
231
+ }
232
+ state = s.policy.OnUnconfirmed(state)
233
+ s.relays[relayURL] = state
234
+}
235
+
236
func (s *RelaySet) ApplyRelayDiscoveryResponse(targetURL string, resp types.DiscoveryResponse, now time.Time) (relaySetChanged bool, err error) {
237
if now.IsZero() {
238
now = time.Now().UTC()
239
} else {
240
now = now.UTC()
241
}
200
-
201
- if resp.ProtocolVersion != types.DiscoveryVersion {
202
- return false, fmt.Errorf("relay discovery protocol version mismatch: relay=%q client=%q", resp.ProtocolVersion, types.DiscoveryVersion)
203
- }
242
+ protocolMismatch := resp.ProtocolVersion != types.DiscoveryVersion
243
authoritative := targetURL != ""
244
245
s.mu.Lock()
246
defer s.mu.Unlock()
247
248
discoveredByURL := make(map[string]RelayState, len(resp.Relays))
210
- discoveredOrder := make([]string, 0, len(resp.Relays))
249
+ discoveredOrder := make([]string, 0, len(resp.Relays)+1)
250
targetFound := false
212
- for _, descriptor := range resp.Relays {
251
+ add := func(descriptor types.RelayDescriptor) {
252
relayState, err := newRelayState(descriptor, now)
253
if err != nil {
215
- continue
254
+ return
255
}
256
relayURL := relayState.Descriptor.APIHTTPSAddr
257
if relayURL == "" {
219
- continue
258
+ return
259
}
260
if authoritative && relayURL == targetURL {
261
targetFound = true
@@ -226,30 +265,29 @@ func (s *RelaySet) ApplyRelayDiscoveryResponse(targetURL string, resp types.Disc
265
}
266
discoveredByURL[relayURL] = relayState
267
}
229
-
230
- if authoritative && !targetFound {
231
- return false, errors.New("target relay descriptor missing from relays")
268
+ for _, descriptor := range resp.Relays {
269
+ add(descriptor)
270
}
271
+ missingTarget := authoritative && !targetFound
272
273
for _, relayURL := range discoveredOrder {
274
record := discoveredByURL[relayURL]
275
existingAtURL, hasExistingAtURL := s.relays[relayURL]
276
record.Bootstrap = record.Bootstrap || existingAtURL.Bootstrap
238
- record.Reachable = record.Reachable || existingAtURL.Reachable
277
record.Confirmed = record.Confirmed || existingAtURL.Confirmed
278
record.Banned = record.Banned || existingAtURL.Banned
279
if record.consecutiveFailures < existingAtURL.consecutiveFailures {
280
record.consecutiveFailures = existingAtURL.consecutiveFailures
281
}
282
+ record.nextDirectRefreshAt = existingAtURL.nextDirectRefreshAt
283
if record.DiscoveryRTTAt.IsZero() || (!existingAtURL.DiscoveryRTTAt.IsZero() && existingAtURL.DiscoveryRTTAt.After(record.DiscoveryRTTAt)) {
284
record.DiscoveryRTT = existingAtURL.DiscoveryRTT
285
record.DiscoveryRTTAt = existingAtURL.DiscoveryRTTAt
286
}
287
249
- if authoritative && relayURL == targetURL {
250
- record = s.policy.OnConfirmed(record)
251
- } else {
252
- record = s.policy.OnHinted(record)
288
+ if !protocolMismatch && !missingTarget && authoritative && relayURL == targetURL {
289
+ record.consecutiveFailures = 0
290
+ record.nextDirectRefreshAt = time.Time{}
291
}
292
293
s.relays[relayURL] = record
@@ -258,6 +296,12 @@ func (s *RelaySet) ApplyRelayDiscoveryResponse(targetURL string, resp types.Disc
296
relaySetChanged = true
297
}
298
}
299
+ if missingTarget {
300
+ return relaySetChanged, errors.New("target relay descriptor missing from relays")
301
+ }
302
+ if protocolMismatch && authoritative {
303
+ return relaySetChanged, fmt.Errorf("relay discovery protocol version mismatch: relay=%q client=%q", resp.ProtocolVersion, types.DiscoveryVersion)
304
+ }
305
return relaySetChanged, nil
306
}
307
@@ -275,7 +319,7 @@ func (s *RelaySet) RecordDiscoveryRTT(relayURL string, rtt time.Duration, measur
319
s.relays[relayURL] = state
320
}
321
278
-func (s *RelaySet) RecordRelayFailure(relayURL string, err error, recoveryFailures int) (expired bool, expireReason string, consecutiveFailures int) {
322
+func (s *RelaySet) RecordRelayFailure(relayURL string, err error, recoveryFailures int) (backedOff bool, backoffReason string, consecutiveFailures int) {
323
s.mu.Lock()
324
defer s.mu.Unlock()
325
@@ -283,7 +327,7 @@ func (s *RelaySet) RecordRelayFailure(relayURL string, err error, recoveryFailur
327
if !ok {
328
return false, "", 0
329
}
286
- state, expired, expireReason = s.policy.OnFailure(state, err, recoveryFailures)
330
+ state, backedOff, backoffReason = s.policy.OnFailure(state, err, recoveryFailures)
331
s.relays[relayURL] = state
288
- return expired, expireReason, state.consecutiveFailures
332
+ return backedOff, backoffReason, state.consecutiveFailures
333
}
portal/discovery/relayset_test.go
+253
-2
@@ -21,11 +21,262 @@ func TestApplyRelayDiscoveryResponsePreservesBootstrapFlag(t *testing.T) {
21
t.Fatalf("ApplyRelayDiscoveryResponse() error = %v", err)
22
}
23
24
- states := set.ActiveRelays()
24
+ states := set.AggregateRelays()
25
if len(states) != 1 {
26
- t.Fatalf("len(ActiveRelays()) = %d, want 1", len(states))
26
+ t.Fatalf("len(AggregateRelays()) = %d, want 1", len(states))
27
}
28
if !states[0].Bootstrap {
29
t.Fatal("bootstrap relay lost bootstrap flag after discovery update")
30
}
31
}
32
+
33
+func TestDescriptorsKeepsStaleRelayHintWithinRetentionWindow(t *testing.T) {
34
+ set, err := NewRelaySet(nil)
35
+ if err != nil {
36
+ t.Fatalf("NewRelaySet() error = %v", err)
37
+ }
38
+
39
+ now := time.Now().UTC()
40
+ relayURL := "https://relay-stale.example"
41
+ state := confirmedPolicyRelayState(t, "relay-stale", relayURL)
42
+ state.Descriptor.ExpiresAt = now.Add(-time.Minute)
43
+ state.LastSeenAt = now.Add(-6 * time.Hour)
44
+ state.Descriptor.SupportsUDP = true
45
+ state.Descriptor.SupportsTCP = true
46
+ state.Descriptor.SupportsOverlayPeer = true
47
+ state.Descriptor.IngressTLSAddr = "relay-stale.example:443"
48
+ state.Descriptor.WireGuardPublicKey = "pub"
49
+ state.Descriptor.WireGuardEndpoint = "relay-stale.example:51820"
50
+ state.Descriptor.OverlayIPv4 = "10.0.0.1"
51
+ state.Descriptor.OverlayCIDRs = []string{"10.0.0.0/24"}
52
+ state.Descriptor.Load = 1
53
+ state.Descriptor.LoadScore = 2
54
+
55
+ set.mu.Lock()
56
+ set.relays[relayURL] = state
57
+ set.mu.Unlock()
58
+
59
+ descriptors := set.Descriptors()
60
+ if len(descriptors) != 1 {
61
+ t.Fatalf("len(Descriptors()) = %d, want 1", len(descriptors))
62
+ }
63
+ got := descriptors[0]
64
+ if got.APIHTTPSAddr != relayURL {
65
+ t.Fatalf("descriptor api_https_addr = %q, want %q", got.APIHTTPSAddr, relayURL)
66
+ }
67
+ if !got.ExpiresAt.After(now) {
68
+ t.Fatalf("descriptor expires_at = %v, want future expiry", got.ExpiresAt)
69
+ }
70
+ if !got.SupportsUDP || !got.SupportsTCP || !got.SupportsOverlayPeer {
71
+ t.Fatal("stale advertised descriptor should preserve last known capability claims")
72
+ }
73
+ if got.IngressTLSAddr == "" || got.WireGuardPublicKey == "" || got.WireGuardEndpoint == "" || got.OverlayIPv4 == "" || len(got.OverlayCIDRs) == 0 {
74
+ t.Fatal("stale advertised descriptor should preserve last known routing fields")
75
+ }
76
+ if got.Load != 1 || got.LoadScore != 2 {
77
+ t.Fatal("stale advertised descriptor should preserve last known load signals")
78
+ }
79
+}
80
+
81
+func TestDescriptorsDropsRelayAfterHintRetentionWindow(t *testing.T) {
82
+ set, err := NewRelaySet(nil)
83
+ if err != nil {
84
+ t.Fatalf("NewRelaySet() error = %v", err)
85
+ }
86
+
87
+ now := time.Now().UTC()
88
+ relayURL := "https://relay-old.example"
89
+ state := confirmedPolicyRelayState(t, "relay-old", relayURL)
90
+ state.Descriptor.ExpiresAt = now.Add(-time.Minute)
91
+ state.LastSeenAt = now.Add(-DiscoveryHintRetentionTTL).Add(-time.Minute)
92
+
93
+ set.mu.Lock()
94
+ set.relays[relayURL] = state
95
+ set.mu.Unlock()
96
+
97
+ descriptors := set.Descriptors()
98
+ if len(descriptors) != 0 {
99
+ t.Fatalf("len(Descriptors()) = %d, want 0", len(descriptors))
100
+ }
101
+}
102
+
103
+func TestApplyRelayDiscoveryResponseCollectsRelaysDespiteProtocolMismatch(t *testing.T) {
104
+ set, err := NewRelaySet(nil)
105
+ if err != nil {
106
+ t.Fatalf("NewRelaySet() error = %v", err)
107
+ }
108
+
109
+ desc := mustPolicyRelayDescriptor(t, "relay-mismatch", "https://relay-mismatch.example")
110
+ changed, err := set.ApplyRelayDiscoveryResponse("", types.DiscoveryResponse{
111
+ ProtocolVersion: "5",
112
+ Relays: []types.RelayDescriptor{desc},
113
+ }, time.Now().UTC())
114
+ if err != nil {
115
+ t.Fatalf("ApplyRelayDiscoveryResponse() error = %v", err)
116
+ }
117
+ if !changed {
118
+ t.Fatal("expected protocol-mismatched discovery response to change relay set")
119
+ }
120
+
121
+ states := set.AggregateRelays()
122
+ if len(states) != 1 {
123
+ t.Fatalf("len(AggregateRelays()) = %d, want 1", len(states))
124
+ }
125
+ if got := states[0].Descriptor.APIHTTPSAddr; got != desc.APIHTTPSAddr {
126
+ t.Fatalf("states[0] = %q, want %q", got, desc.APIHTTPSAddr)
127
+ }
128
+ if states[0].Confirmed {
129
+ t.Fatal("hinted relay should not become locally confirmed from aggregation")
130
+ }
131
+}
132
+
133
+func TestApplyRelayDiscoveryResponseCollectsHintsWhenTargetDescriptorIsMissing(t *testing.T) {
134
+ set, err := NewRelaySet(nil)
135
+ if err != nil {
136
+ t.Fatalf("NewRelaySet() error = %v", err)
137
+ }
138
+
139
+ hinted := mustPolicyRelayDescriptor(t, "relay-hinted", "https://relay-hinted.example")
140
+ changed, err := set.ApplyRelayDiscoveryResponse("https://relay-source.example", types.DiscoveryResponse{
141
+ ProtocolVersion: "5",
142
+ Relays: []types.RelayDescriptor{hinted},
143
+ }, time.Now().UTC())
144
+ if err == nil {
145
+ t.Fatal("expected missing target descriptor error")
146
+ }
147
+ if !changed {
148
+ t.Fatal("expected hinted relay to still be collected")
149
+ }
150
+
151
+ states := set.AggregateRelays()
152
+ if len(states) != 1 {
153
+ t.Fatalf("len(AggregateRelays()) = %d, want 1", len(states))
154
+ }
155
+ if got := states[0].Descriptor.APIHTTPSAddr; got != hinted.APIHTTPSAddr {
156
+ t.Fatalf("states[0] = %q, want %q", got, hinted.APIHTTPSAddr)
157
+ }
158
+ if states[0].Confirmed {
159
+ t.Fatal("hinted relay should not become locally confirmed when target descriptor is missing")
160
+ }
161
+}
162
+
163
+func TestApplyRelayDiscoveryResponseClearsDirectRetryOnAuthoritativeSuccess(t *testing.T) {
164
+ set, err := NewRelaySet(nil)
165
+ if err != nil {
166
+ t.Fatalf("NewRelaySet() error = %v", err)
167
+ }
168
+
169
+ relayURL := "https://relay-source.example"
170
+ desc := mustPolicyRelayDescriptor(t, "relay-source", relayURL)
171
+ set.mu.Lock()
172
+ state := RelayState{
173
+ Descriptor: desc,
174
+ LastSeenAt: time.Now().UTC(),
175
+ consecutiveFailures: defaultRecoveryFailures,
176
+ nextDirectRefreshAt: time.Now().UTC().Add(time.Minute),
177
+ }
178
+ set.relays[relayURL] = state
179
+ set.mu.Unlock()
180
+
181
+ if _, err := set.ApplyRelayDiscoveryResponse(relayURL, types.DiscoveryResponse{
182
+ ProtocolVersion: types.DiscoveryVersion,
183
+ Relays: []types.RelayDescriptor{desc},
184
+ }, time.Now().UTC()); err != nil {
185
+ t.Fatalf("ApplyRelayDiscoveryResponse() error = %v", err)
186
+ }
187
+
188
+ set.mu.RLock()
189
+ refreshed := set.relays[relayURL]
190
+ set.mu.RUnlock()
191
+ if refreshed.consecutiveFailures != 0 {
192
+ t.Fatalf("consecutiveFailures = %d, want 0", refreshed.consecutiveFailures)
193
+ }
194
+ if !refreshed.nextDirectRefreshAt.IsZero() {
195
+ t.Fatalf("nextDirectRefreshAt = %v, want zero time", refreshed.nextDirectRefreshAt)
196
+ }
197
+}
198
+
199
+func TestApplyRelayDiscoveryResponsePreservesDirectRetryOnHint(t *testing.T) {
200
+ set, err := NewRelaySet(nil)
201
+ if err != nil {
202
+ t.Fatalf("NewRelaySet() error = %v", err)
203
+ }
204
+
205
+ relayURL := "https://relay-hinted.example"
206
+ desc := mustPolicyRelayDescriptor(t, "relay-hinted", relayURL)
207
+ nextDirectRefreshAt := time.Now().UTC().Add(time.Minute)
208
+ set.mu.Lock()
209
+ state := RelayState{
210
+ Descriptor: desc,
211
+ LastSeenAt: time.Now().UTC(),
212
+ nextDirectRefreshAt: nextDirectRefreshAt,
213
+ }
214
+ set.relays[relayURL] = state
215
+ set.mu.Unlock()
216
+
217
+ if _, err := set.ApplyRelayDiscoveryResponse("", types.DiscoveryResponse{
218
+ ProtocolVersion: types.DiscoveryVersion,
219
+ Relays: []types.RelayDescriptor{desc},
220
+ }, time.Now().UTC()); err != nil {
221
+ t.Fatalf("ApplyRelayDiscoveryResponse() error = %v", err)
222
+ }
223
+
224
+ set.mu.RLock()
225
+ refreshed := set.relays[relayURL]
226
+ set.mu.RUnlock()
227
+ if !refreshed.nextDirectRefreshAt.Equal(nextDirectRefreshAt) {
228
+ t.Fatalf("nextDirectRefreshAt = %v, want %v", refreshed.nextDirectRefreshAt, nextDirectRefreshAt)
229
+ }
230
+}
231
+
232
+func TestConfirmRelayURLMarksRelayConfirmedWithoutChangingAggregateDescriptor(t *testing.T) {
233
+ set, err := NewRelaySet(nil)
234
+ if err != nil {
235
+ t.Fatalf("NewRelaySet() error = %v", err)
236
+ }
237
+
238
+ relayURL := "https://relay-confirmed.example"
239
+ state := RelayState{
240
+ Descriptor: mustPolicyRelayDescriptor(t, "relay-confirmed", relayURL),
241
+ LastSeenAt: time.Now().UTC(),
242
+ }
243
+
244
+ set.mu.Lock()
245
+ set.relays[relayURL] = state
246
+ set.mu.Unlock()
247
+
248
+ set.ConfirmRelayURL(relayURL)
249
+
250
+ set.mu.RLock()
251
+ confirmed := set.relays[relayURL]
252
+ set.mu.RUnlock()
253
+ if !confirmed.Confirmed {
254
+ t.Fatal("relay should become locally confirmed after listener success")
255
+ }
256
+ if confirmed.Descriptor.APIHTTPSAddr != relayURL {
257
+ t.Fatalf("descriptor api_https_addr = %q, want %q", confirmed.Descriptor.APIHTTPSAddr, relayURL)
258
+ }
259
+}
260
+
261
+func TestUnconfirmRelayURLClearsLocalConfirmationOnly(t *testing.T) {
262
+ set, err := NewRelaySet(nil)
263
+ if err != nil {
264
+ t.Fatalf("NewRelaySet() error = %v", err)
265
+ }
266
+
267
+ relayURL := "https://relay-confirmed.example"
268
+ state := confirmedPolicyRelayState(t, "relay-confirmed", relayURL)
269
+
270
+ set.mu.Lock()
271
+ set.relays[relayURL] = state
272
+ set.mu.Unlock()
273
+
274
+ set.UnconfirmRelayURL(relayURL)
275
+
276
+ set.mu.RLock()
277
+ unconfirmed := set.relays[relayURL]
278
+ set.mu.RUnlock()
279
+ if unconfirmed.Confirmed {
280
+ t.Fatal("relay should lose local confirmation after listener failure")
281
+ }
282
+}
portal/discovery/relaystate.go
+14
-22
@@ -8,17 +8,25 @@ import (
8
"github.com/gosuda/portal-tunnel/v2/utils"
9
)
10
11
+const (
12
+ DiscoveryDescriptorTTL = 5 * time.Minute
13
+ DiscoveryHintRetentionTTL = 30 * 24 * time.Hour
14
+ defaultDirectRecoveryBackoff = 1 * time.Minute
15
+ maxDirectRecoveryBackoff = 5 * time.Minute
16
+)
17
+
18
type RelayState struct {
12
- Descriptor types.RelayDescriptor
13
- Bootstrap bool
14
- Reachable bool
15
- Confirmed bool
16
- Banned bool
17
- LastSeenAt time.Time
19
+ Descriptor types.RelayDescriptor
20
+ Bootstrap bool
21
+ Confirmed bool
22
+ Banned bool
23
+ LastSeenAt time.Time
24
+
25
DiscoveryRTT time.Duration
26
DiscoveryRTTAt time.Time
27
28
consecutiveFailures int
29
+ nextDirectRefreshAt time.Time
30
}
31
32
type ClientState struct {
@@ -66,19 +74,3 @@ func newRelayStateFromURL(relayURL string) RelayState {
74
func (state RelayState) hasDescriptor() bool {
75
return !state.LastSeenAt.IsZero()
76
}
69
-
70
-func (state RelayState) discoverable(now time.Time) bool {
71
- if state.Banned {
72
- return false
73
- }
74
- if !state.hasDescriptor() {
75
- return state.Bootstrap
76
- }
77
- if !state.Bootstrap && !state.Reachable {
78
- return false
79
- }
80
- if !state.Descriptor.ExpiresAt.After(now) {
81
- return false
82
- }
83
- return true
84
-}
sdk/expose_test.go
-36
@@ -71,42 +71,6 @@ func TestExposureReconcileRemovesBannedRelayFromActiveSet(t *testing.T) {
71
}
72
}
73
74
-func TestExposureReconcileSkipsBannedRelay(t *testing.T) {
75
- const (
76
- relayA = "https://relay-a.example"
77
- relayB = "https://relay-b.example"
78
- )
79
-
80
- exposure := &Exposure{
81
- relaySet: mustRelaySet(t),
82
- relayListeners: make(map[string]*Listener, 1),
83
- }
84
- exposure.relaySet.BanRelayURL(relayB)
85
- exposure.relayListeners = map[string]*Listener{
86
- relayA: {},
87
- }
88
-
89
- if err := exposure.relaySet.SetBootstrapRelayURLs([]string{relayA, relayB}); err != nil {
90
- t.Fatalf("SetBootstrapRelayURLs() error = %v", err)
91
- }
92
- if err := exposure.reconcileRelayListeners(false); err != nil {
93
- t.Fatalf("reconcileRelayListeners() error = %v", err)
94
- }
95
- if got := exposure.ActiveRelayURLs(); len(got) != 1 || got[0] != relayA {
96
- t.Fatalf("ActiveRelayURLs() = %v, want [%q]", got, relayA)
97
- }
98
- exposure.listenerMu.RLock()
99
- _, relayAExists := exposure.relayListeners[relayA]
100
- _, relayBExists := exposure.relayListeners[relayB]
101
- exposure.listenerMu.RUnlock()
102
- if !relayAExists {
103
- t.Fatal("active relay listener missing from exposure.listeners")
104
- }
105
- if relayBExists {
106
- t.Fatal("banned relay listener should not be added to exposure.listeners")
107
- }
108
-}
109
-
74
func TestExposureReconcileRemovesStaleListener(t *testing.T) {
75
const (
76
relayA = "https://relay-a.example"
sdk/listener.go
+6
@@ -148,6 +148,7 @@ func (l *Listener) runStartup(ctx context.Context) {
148
errors.Is(err, &types.APIRequestError{Code: types.APIErrorCodeHostnameConflict}) ||
149
errors.Is(err, &types.APIRequestError{Code: types.APIErrorCodeIPBanned}) {
150
if l.relaySet != nil && l.api != nil && l.api.baseURL != nil {
151
+ l.relaySet.UnconfirmRelayURL(l.api.baseURL.String())
152
l.relaySet.RecordRelayFailure(l.api.baseURL.String(), err, 1)
153
}
154
log.Error().
@@ -656,6 +657,9 @@ func (l *Listener) registerAndConfigure(ctx context.Context) error {
657
if datagram != nil {
658
datagram.Clear("lease updated")
659
}
660
+ if l.relaySet != nil && l.api != nil && l.api.baseURL != nil {
661
+ l.relaySet.ConfirmRelayURL(l.api.baseURL.String())
662
+ }
663
l.registerOnce.Do(func() { close(l.registered) })
664
return nil
665
}
@@ -673,6 +677,7 @@ func (l *Listener) retryOrClose(ctx context.Context, operation string, err error
677
678
if l.retryCount > 0 && retries > l.retryCount {
679
if l.relaySet != nil && l.api != nil && l.api.baseURL != nil {
680
+ l.relaySet.UnconfirmRelayURL(l.api.baseURL.String())
681
l.relaySet.RecordRelayFailure(l.api.baseURL.String(), err, 1)
682
}
683
if operation != "lease renewal" {
@@ -729,6 +734,7 @@ func (l *Listener) closed() bool {
734
735
func (l *Listener) ban() {
736
if l.relaySet != nil && l.api != nil && l.api.baseURL != nil {
737
+ l.relaySet.UnconfirmRelayURL(l.api.baseURL.String())
738
l.relaySet.BanRelayURL(l.api.baseURL.String())
739
}
740
_ = l.Close()
sdk/mitm_test.go
+2
-2
@@ -228,7 +228,7 @@ func TestMITMProbeDetectionBansListener(t *testing.T) {
228
Reason: types.MITMProbeReasonExporterMismatch,
229
}, nil)
230
231
- for _, state := range listener.relaySet.ActiveRelays() {
231
+ for _, state := range listener.relaySet.AggregateRelays() {
232
if state.Descriptor.APIHTTPSAddr == relayURL.String() {
233
t.Fatal("relay still active after mitm detection")
234
}
@@ -261,7 +261,7 @@ func TestMITMProbeDetectionWarnsWithoutBanningListener(t *testing.T) {
261
}, nil)
262
263
activeRelayURLs := make([]string, 0)
264
- for _, state := range listener.relaySet.ActiveRelays() {
264
+ for _, state := range listener.relaySet.AggregateRelays() {
265
activeRelayURLs = append(activeRelayURLs, state.Descriptor.APIHTTPSAddr)
266
}
267
if len(activeRelayURLs) != 1 || activeRelayURLs[0] != relayURL.String() {