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() {