Refactor relay selection and routing logic

Kim committed May 19, 2026 at 15:45 UTC 19480146c96048962a5f7f37520ea4962af583cf
13 files changed +349 -359
portal/discovery/announce_test.go
+6 -6
@@ -53,8 +53,8 @@ func TestInsertAnnouncedAcceptsValidDescriptor(t *testing.T) {
53 if err := set.InsertAnnounced(desc, now); err != nil {
54 t.Fatalf("InsertAnnounced() error = %v", err)
55 }
56 - if got := set.AggregateRelays(); len(got) != 1 {
57 - t.Fatalf("len(AggregateRelays()) = %d, want 1", len(got))
56 + if got := relayStates(set); len(got) != 1 {
57 + t.Fatalf("len(relayStates()) = %d, want 1", len(got))
58 }
59 }
60
@@ -82,9 +82,9 @@ func TestInsertAnnouncedIgnoresSupersededRollback(t *testing.T) {
82 t.Fatalf("superseded insert error = %v", err)
83 }
84
85 - states := set.AggregateRelays()
85 + states := relayStates(set)
86 if len(states) != 1 {
87 - t.Fatalf("len(AggregateRelays()) = %d, want 1", len(states))
87 + t.Fatalf("len(relayStates()) = %d, want 1", len(states))
88 }
89 if got := states[0].Descriptor.IssuedAt; !got.Equal(newer.IssuedAt) {
90 t.Fatalf("stored issued_at = %v, want %v", got, newer.IssuedAt)
@@ -122,9 +122,9 @@ func TestInsertAnnouncedBlocksCrossIdentityTakeover(t *testing.T) {
122 t.Fatal("expected takeover reject")
123 }
124
125 - states := set.AggregateRelays()
125 + states := relayStates(set)
126 if len(states) != 1 {
127 - t.Fatalf("len(AggregateRelays()) = %d, want 1", len(states))
127 + t.Fatalf("len(relayStates()) = %d, want 1", len(states))
128 }
129 if got := states[0].Descriptor.Address; got != owner.Address {
130 t.Fatalf("retained address = %q, want %q", got, owner.Address)
portal/discovery/mols.go
+35 -112
@@ -1,6 +1,6 @@
1 package discovery
2
3 -// MOLSRelayPolicy uses a GF(64) MOLS-derived score as the primary
3 +// MOLS route selection uses a GF(64) MOLS-derived score as the primary
4 // deterministic ordering for eligible relays. Health and freshness gates decide
5 // eligibility before the MOLS score is applied.
6
@@ -131,86 +131,7 @@ func molsRTTStats(states []RelayState) (mean time.Duration, cv float64) {
131 return time.Duration(avg), cv
132 }
133
134 -func isRelayFallback(state RelayState) bool {
135 - return !state.DiscoveryRTTAt.IsZero() && state.DiscoveryRTT > molsFallbackRTTThreshold
136 -}
137 -
138 -type MOLSRelayPolicy struct{}
139 -
140 -func (p MOLSRelayPolicy) SelectAggregate(states []RelayState) []RelayState {
141 - out := make([]RelayState, 0, len(states))
142 - for _, s := range states {
143 - if !s.Banned {
144 - out = append(out, s)
145 - }
146 - }
147 - return out
148 -}
149 -
150 -func (p MOLSRelayPolicy) SelectConfirmed(states []RelayState) []RelayState {
151 - out := make([]RelayState, 0)
152 - for _, s := range states {
153 - if s.Confirmed {
154 - out = append(out, s)
155 - }
156 - }
157 - return out
158 -}
159 -
160 -func (p MOLSRelayPolicy) OnActiveConfirmed(state RelayState) RelayState {
161 - state.Confirmed = true
162 - state.activeFailures = 0
163 - state.suppressActiveUntil = time.Time{}
164 - return state
165 -}
166 -
167 -func (p MOLSRelayPolicy) OnUnconfirmed(state RelayState) RelayState {
168 - state.Confirmed = false
169 - return state
170 -}
171 -
172 -func (p MOLSRelayPolicy) OnDiscoveryConfirmed(state RelayState) RelayState {
173 - state.discoveryFailures = 0
174 - state.nextDiscoveryRefreshAt = time.Time{}
175 - return state
176 -}
177 -
178 -func (p MOLSRelayPolicy) OnDiscoveryFailure(state RelayState, err error, recoveryFailures int) (RelayState, bool, string) {
179 - state.discoveryFailures++
180 -
181 - if recoveryFailures <= 0 || state.discoveryFailures < recoveryFailures {
182 - return state, false, "retry"
183 - }
184 - failuresOverBudget := state.discoveryFailures - recoveryFailures
185 - backoff := defaultDirectRecoveryBackoff << min(failuresOverBudget, 3)
186 - if backoff > maxDirectRecoveryBackoff {
187 - backoff = maxDirectRecoveryBackoff
188 - }
189 - state.nextDiscoveryRefreshAt = time.Now().Add(backoff)
190 - return state, true, "discovery"
191 -}
192 -
193 -func (p MOLSRelayPolicy) OnActiveFailure(state RelayState, err error, recoveryFailures int) (RelayState, bool, string) {
194 - state.activeFailures++
195 -
196 - if recoveryFailures <= 0 || state.activeFailures < recoveryFailures {
197 - return state, false, "retry"
198 - }
199 - failuresOverBudget := state.activeFailures - recoveryFailures
200 - backoff := defaultDirectRecoveryBackoff << min(failuresOverBudget, 3)
201 - if backoff > maxDirectRecoveryBackoff {
202 - backoff = maxDirectRecoveryBackoff
203 - }
204 - state.suppressActiveUntil = time.Now().Add(backoff)
205 - return state, true, "active"
206 -}
207 -
208 -func (p MOLSRelayPolicy) OnBanned(state RelayState) RelayState {
209 - state.Banned = true
210 - return state
211 -}
212 -
213 -func (p MOLSRelayPolicy) rankRelayPool(autoPool []RelayState, localAddress string) []string {
134 +func rankRelayPool(autoPool []RelayState, localAddress string) []string {
135 if len(autoPool) == 0 {
136 return nil
137 }
@@ -228,7 +149,7 @@ func (p MOLSRelayPolicy) rankRelayPool(autoPool []RelayState, localAddress strin
149 active := make([]RelayState, 0, len(autoPool))
150 fallbacks := make([]RelayState, 0)
151 for _, state := range autoPool {
231 - if isRelayFallback(state) {
152 + if !state.DiscoveryRTTAt.IsZero() && state.DiscoveryRTT > molsFallbackRTTThreshold {
153 fallbacks = append(fallbacks, state)
154 } else {
155 active = append(active, state)
@@ -297,23 +218,21 @@ func (p MOLSRelayPolicy) rankRelayPool(autoPool []RelayState, localAddress strin
218 return autoURLs
219 }
220
300 -func (p MOLSRelayPolicy) SelectPriority(states []RelayState, clientState ClientState) []string {
301 - selected := p.SelectAggregate(states)
302 - if len(selected) == 0 {
303 - return nil
304 - }
305 -
221 +func selectPriority(states []RelayState, routeState RouteState) []string {
222 now := time.Now().UTC()
223 explicit := make([]string, 0)
308 - autoPool := make([]RelayState, 0, len(selected))
309 - for _, state := range selected {
224 + autoPool := make([]RelayState, 0, len(states))
225 + for _, state := range states {
226 + if state.Banned {
227 + continue
228 + }
229 relayURL := state.Descriptor.APIHTTPSAddr
311 - if slices.Contains(clientState.ExplicitRelayURLs, relayURL) {
230 + if slices.Contains(routeState.ExplicitRelayURLs, relayURL) {
231 if state.hasObservedDescriptor() && state.Descriptor.ExpiresAt.After(now) {
313 - if clientState.RequireUDP && !state.Descriptor.SupportsUDP {
232 + if routeState.RequireUDP && !state.Descriptor.SupportsUDP {
233 continue
234 }
316 - if clientState.RequireTCP && !state.Descriptor.SupportsTCP {
235 + if routeState.RequireTCP && !state.Descriptor.SupportsTCP {
236 continue
237 }
238 }
@@ -325,10 +244,10 @@ func (p MOLSRelayPolicy) SelectPriority(states []RelayState, clientState ClientS
244 if !state.Descriptor.ExpiresAt.After(now) {
245 continue
246 }
328 - if clientState.RequireUDP && !state.Descriptor.SupportsUDP {
247 + if routeState.RequireUDP && !state.Descriptor.SupportsUDP {
248 continue
249 }
331 - if clientState.RequireTCP && !state.Descriptor.SupportsTCP {
250 + if routeState.RequireTCP && !state.Descriptor.SupportsTCP {
251 continue
252 }
253 }
@@ -338,8 +257,8 @@ func (p MOLSRelayPolicy) SelectPriority(states []RelayState, clientState ClientS
257 autoPool = append(autoPool, state)
258 }
259
341 - autoURLs := p.rankRelayPool(autoPool, clientState.LocalAddress)
342 - maxActiveRelays := clientState.MaxActiveRelays
260 + autoURLs := rankRelayPool(autoPool, routeState.LocalAddress)
261 + maxActiveRelays := routeState.MaxActiveRelays
262 if maxActiveRelays <= 0 {
263 maxActiveRelays = defaultMaxActiveRelays
264 }
@@ -349,26 +268,30 @@ func (p MOLSRelayPolicy) SelectPriority(states []RelayState, clientState ClientS
268 return append(explicit, autoURLs...)
269 }
270
352 -func (p MOLSRelayPolicy) SelectMultiHop(states []RelayState, clientState ClientState) []string {
353 - if clientState.MultiHopDepth <= 1 {
354 - return nil
355 - }
356 -
357 - selected := p.SelectAggregate(states)
358 - if len(selected) == 0 {
271 +func selectMultiHop(states []RelayState, routeState RouteState) []string {
272 + if routeState.MultiHopDepth <= 1 {
273 return nil
274 }
275
276 now := time.Now().UTC()
363 - autoPool := make([]RelayState, 0, len(selected))
364 - for _, state := range selected {
365 - if clientState.RequireUDP && state.hasObservedDescriptor() && !state.Descriptor.SupportsUDP {
277 + autoPool := make([]RelayState, 0, len(states))
278 + for _, state := range states {
279 + if state.Banned {
280 + continue
281 + }
282 + if !state.hasObservedDescriptor() {
283 + continue
284 + }
285 + if routeState.RequireUDP && !state.Descriptor.SupportsUDP {
286 + continue
287 + }
288 + if routeState.RequireTCP && !state.Descriptor.SupportsTCP {
289 continue
290 }
368 - if clientState.RequireTCP && state.hasObservedDescriptor() && !state.Descriptor.SupportsTCP {
291 + if !state.Descriptor.ExpiresAt.After(now) {
292 continue
293 }
371 - if !state.hasObservedDescriptor() || !state.Descriptor.ExpiresAt.After(now) || !state.Descriptor.HasOverlayPeer() {
294 + if !state.Descriptor.HasOverlayPeer() {
295 continue
296 }
297 if !state.suppressActiveUntil.IsZero() && state.suppressActiveUntil.After(now) {
@@ -377,9 +300,9 @@ func (p MOLSRelayPolicy) SelectMultiHop(states []RelayState, clientState ClientS
300 autoPool = append(autoPool, state)
301 }
302
380 - multiHop := p.rankRelayPool(autoPool, clientState.LocalAddress)
381 - if len(multiHop) > clientState.MultiHopDepth {
382 - multiHop = multiHop[:clientState.MultiHopDepth]
303 + multiHop := rankRelayPool(autoPool, routeState.LocalAddress)
304 + if len(multiHop) > routeState.MultiHopDepth {
305 + multiHop = multiHop[:routeState.MultiHopDepth]
306 }
307 return multiHop
308 }
portal/discovery/mols_test.go
+39 -69
@@ -66,7 +66,7 @@ func TestMOLSScoreRange(t *testing.T) {
66 }
67
68 // TestMOLSScoreRowPermutation checks that each row of the MOLS score grid is a
69 -// permutation of 1..n². Rows are indexed by ingress i; columns by candidate j.
69 +// permutation of 1..n^2. Rows are indexed by ingress i; columns by candidate j.
70 func TestMOLSScoreRowPermutation(t *testing.T) {
71 for i := range uint8(64) {
72 seen := make(map[int]struct{}, 64)
@@ -92,7 +92,7 @@ func TestMOLSCongestionScoreRange(t *testing.T) {
92 if s < 1 || s > molsOrder*molsOrder {
93 t.Fatalf("molsCongestionScore(%d, %d) = %d, out of range", i, j, s)
94 }
95 - // Verify B(i,j) = (n²+1) - A(i, n-1-j)
95 + // Verify B(i,j) = (n^2+1) - A(i, n-1-j)
96 want := molsMagicConstant - molsScore(i, (molsOrder-1)-j, molsBaseM1, molsBaseM2)
97 if s != want {
98 t.Fatalf("molsCongestionScore(%d, %d) = %d, want %d", i, j, s, want)
@@ -145,7 +145,7 @@ func TestMOLSRTTStatsCVHigh(t *testing.T) {
145 func TestMOLSRTTStatsSkipsMissingRTT(t *testing.T) {
146 states := []RelayState{
147 {DiscoveryRTT: 100 * time.Millisecond, DiscoveryRTTAt: time.Now()},
148 - {DiscoveryRTT: 999 * time.Second}, // no DiscoveryRTTAt → excluded
148 + {DiscoveryRTT: 999 * time.Second}, // no DiscoveryRTTAt, excluded
149 }
150 mean, _ := molsRTTStats(states)
151 if mean != 100*time.Millisecond {
@@ -153,43 +153,18 @@ func TestMOLSRTTStatsSkipsMissingRTT(t *testing.T) {
153 }
154 }
155
156 -// TestIsRelayFallbackHighRTT checks that a relay with RTT > threshold is
157 -// classified as Fallback.
158 -func TestIsRelayFallbackHighRTT(t *testing.T) {
159 - state := RelayState{
160 - DiscoveryRTT: molsFallbackRTTThreshold + time.Millisecond,
161 - DiscoveryRTTAt: time.Now(),
162 - }
163 - if !isRelayFallback(state) {
164 - t.Fatal("expected high-RTT relay to be classified as Fallback")
165 - }
166 -}
167 -
168 -// TestIsRelayFallbackNormalRTT checks that a relay with normal RTT is not
169 -// classified as Fallback.
170 -func TestIsRelayFallbackNormalRTT(t *testing.T) {
171 - state := RelayState{
172 - DiscoveryRTT: 200 * time.Millisecond,
173 - DiscoveryRTTAt: time.Now(),
174 - }
175 - if isRelayFallback(state) {
176 - t.Fatal("expected normal-RTT relay not to be classified as Fallback")
177 - }
178 -}
179 -
156 // TestMOLSSelectPriorityKeepsExplicitRelaysOutsideAutoLimit verifies that
157 // explicit relays are always included, outside of MaxActiveRelays.
158 func TestMOLSSelectPriorityKeepsExplicitRelaysOutsideAutoLimit(t *testing.T) {
183 - policy := MOLSRelayPolicy{}
159 explicitRelay := "https://relay-explicit.example"
160 relayA := "https://relay-a.example"
161 relayB := "https://relay-b.example"
162
188 - selected := policy.SelectPriority([]RelayState{
163 + selected := selectPriority([]RelayState{
164 bootstrapPolicyRelayState(explicitRelay),
165 confirmedPolicyRelayState(t, relayA),
166 confirmedPolicyRelayState(t, relayB),
192 - }, ClientState{
167 + }, RouteState{
168 ExplicitRelayURLs: []string{explicitRelay},
169 MaxActiveRelays: 1,
170 })
@@ -205,17 +180,16 @@ func TestMOLSSelectPriorityKeepsExplicitRelaysOutsideAutoLimit(t *testing.T) {
180 // TestMOLSSelectPriorityDeterministic verifies that the same inputs always
181 // produce the same ordered output.
182 func TestMOLSSelectPriorityDeterministic(t *testing.T) {
208 - policy := MOLSRelayPolicy{}
183 states := []RelayState{
184 confirmedPolicyRelayState(t, "https://relay-a.example"),
185 confirmedPolicyRelayState(t, "https://relay-b.example"),
186 confirmedPolicyRelayState(t, "https://relay-c.example"),
187 }
214 - clientState := ClientState{LocalAddress: "0x1234abcd"}
188 + routeState := RouteState{LocalAddress: "0x1234abcd"}
189
216 - first := policy.SelectPriority(states, clientState)
190 + first := selectPriority(states, routeState)
191 for range 5 {
218 - got := policy.SelectPriority(states, clientState)
192 + got := selectPriority(states, routeState)
193 if len(got) != len(first) {
194 t.Fatalf("non-deterministic length: %d vs %d", len(got), len(first))
195 }
@@ -230,7 +204,6 @@ func TestMOLSSelectPriorityDeterministic(t *testing.T) {
204 // TestMOLSSelectPriorityFallbackRelaysDemoted checks that relays with high
205 // RTT are placed after healthy relays in the priority list.
206 func TestMOLSSelectPriorityFallbackRelaysDemoted(t *testing.T) {
233 - policy := MOLSRelayPolicy{}
207
208 // Two healthy relays ensure molsMinActiveNodes is met without promoting fallbacks.
209 healthy1 := confirmedPolicyRelayState(t, "https://relay-healthy-1.example")
@@ -245,7 +218,7 @@ func TestMOLSSelectPriorityFallbackRelaysDemoted(t *testing.T) {
218 fallback.DiscoveryRTT = molsFallbackRTTThreshold + time.Millisecond
219 fallback.DiscoveryRTTAt = time.Now()
220
248 - selected := policy.SelectPriority([]RelayState{fallback, healthy1, healthy2}, ClientState{})
221 + selected := selectPriority([]RelayState{fallback, healthy1, healthy2}, RouteState{})
222
223 if len(selected) != 3 {
224 t.Fatalf("len(selected) = %d, want 3", len(selected))
@@ -260,7 +233,6 @@ func TestMOLSSelectPriorityFallbackRelaysDemoted(t *testing.T) {
233 // are fewer than molsMinActiveNodes healthy relays the engine promotes fallback
234 // relays to maintain the minimum.
235 func TestMOLSSelectPriorityMinActiveNodesPromotesFallback(t *testing.T) {
263 - policy := MOLSRelayPolicy{}
236
237 fallback1 := confirmedPolicyRelayState(t, "https://relay-fallback-1.example")
238 fallback1.DiscoveryRTT = molsFallbackRTTThreshold + time.Millisecond
@@ -269,7 +241,7 @@ func TestMOLSSelectPriorityMinActiveNodesPromotesFallback(t *testing.T) {
241 fallback2.DiscoveryRTT = molsFallbackRTTThreshold + time.Millisecond
242 fallback2.DiscoveryRTTAt = time.Now()
243
272 - selected := policy.SelectPriority([]RelayState{fallback1, fallback2}, ClientState{})
244 + selected := selectPriority([]RelayState{fallback1, fallback2}, RouteState{})
245
246 // Both fallbacks should be promoted to meet the minimum of 2.
247 if len(selected) != 2 {
@@ -281,14 +253,13 @@ func TestMOLSSelectPriorityMinActiveNodesPromotesFallback(t *testing.T) {
253 // Reverse-Siamese mode (triggered by high average RTT) produces a different
254 // ordering than normal mode for the same relay set.
255 func TestMOLSSelectPriorityCongestionSwitchChangesOrder(t *testing.T) {
284 - policy := MOLSRelayPolicy{}
256
257 // Two relays with different MOLS column indices so their scores differ.
258 r1 := confirmedPolicyRelayState(t, "https://relay-one.example")
259 r2 := confirmedPolicyRelayState(t, "https://relay-two.example")
260
290 - // Normal mode: no RTT measurements → no congestion.
291 - normal := policy.SelectPriority([]RelayState{r1, r2}, ClientState{
261 + // Normal mode: no RTT measurements, no congestion.
262 + normal := selectPriority([]RelayState{r1, r2}, RouteState{
263 LocalAddress: "ingress-test",
264 })
265
@@ -301,7 +272,7 @@ func TestMOLSSelectPriorityCongestionSwitchChangesOrder(t *testing.T) {
272 r2c.DiscoveryRTT = rttHigh
273 r2c.DiscoveryRTTAt = time.Now()
274
304 - congested := policy.SelectPriority([]RelayState{r1c, r2c}, ClientState{
275 + congested := selectPriority([]RelayState{r1c, r2c}, RouteState{
276 LocalAddress: "ingress-test",
277 })
278
@@ -323,7 +294,7 @@ func TestMOLSSelectPriorityCongestionSwitchChangesOrder(t *testing.T) {
294 if (normal1 > normal2) != (cong1 > cong2) {
295 t.Fatal("expected congestion switch to invert ordering but result matched normal mode")
296 }
326 - // If ordering is the same it means the math happens to agree — acceptable.
297 + // If ordering is the same it means the math happens to agree; acceptable.
298 }
299 }
300
@@ -331,13 +302,12 @@ func TestMOLSSelectPriorityCongestionSwitchChangesOrder(t *testing.T) {
302 // coefficient of variation triggers the variant multipliers (7, 11) rather than
303 // the base (3, 5), producing a different relay ordering from the base grid.
304 func TestMOLSSelectPriorityVariantGridActivatesOnHighCV(t *testing.T) {
334 - policy := MOLSRelayPolicy{}
305
306 r1 := confirmedPolicyRelayState(t, "https://relay-one.example")
307 r2 := confirmedPolicyRelayState(t, "https://relay-two.example")
308
339 - // Normal mode (no RTT → no congestion, no CV).
340 - normalOrder := policy.SelectPriority([]RelayState{r1, r2}, ClientState{
309 + // Normal mode: no RTT, no congestion, no CV.
310 + normalOrder := selectPriority([]RelayState{r1, r2}, RouteState{
311 LocalAddress: "ingress-cv",
312 })
313
@@ -355,7 +325,7 @@ func TestMOLSSelectPriorityVariantGridActivatesOnHighCV(t *testing.T) {
325 t.Fatalf("test precondition: cv = %v, want > %v", cv, molsCVThreshold)
326 }
327
358 - variantOrder := policy.SelectPriority([]RelayState{r1v, r2v}, ClientState{
328 + variantOrder := selectPriority([]RelayState{r1v, r2v}, RouteState{
329 LocalAddress: "ingress-cv",
330 })
331
@@ -390,7 +360,6 @@ func TestMOLSSelectPriorityVariantGridActivatesOnHighCV(t *testing.T) {
360 // different ingress identities can produce different relay orderings (MOLS
361 // property: each row is an independent permutation).
362 func TestMOLSSelectPriorityDifferentIngressDifferentOrder(t *testing.T) {
393 - policy := MOLSRelayPolicy{}
363
364 r1 := confirmedPolicyRelayState(t, "https://relay-alpha.example")
365 r2 := confirmedPolicyRelayState(t, "https://relay-beta.example")
@@ -404,7 +373,7 @@ func TestMOLSSelectPriorityDifferentIngressDifferentOrder(t *testing.T) {
373 "0xabc", "0xdef", "0x123", "0x456", "user@example.com", "relay.net",
374 }
375 for _, addr := range addresses {
407 - sel := policy.SelectPriority(states, ClientState{LocalAddress: addr})
376 + sel := selectPriority(states, RouteState{LocalAddress: addr})
377 key := ""
378 for _, u := range sel {
379 key += u + "|"
@@ -438,8 +407,7 @@ func TestMOLSSelectPriorityDifferentIngressDifferentOrder(t *testing.T) {
407
408 // TestMOLSSelectPriorityEmptyPoolReturnsNil checks the empty-input guard.
409 func TestMOLSSelectPriorityEmptyPoolReturnsNil(t *testing.T) {
441 - policy := MOLSRelayPolicy{}
442 - if got := policy.SelectPriority(nil, ClientState{}); got != nil {
410 + if got := selectPriority(nil, RouteState{}); got != nil {
411 t.Fatalf("SelectPriority(nil, ...) = %v, want nil", got)
412 }
413 }
@@ -447,50 +415,55 @@ func TestMOLSSelectPriorityEmptyPoolReturnsNil(t *testing.T) {
415 // TestMOLSSelectPriorityMaxActiveRelaysLimitsAutoPool ensures that
416 // MaxActiveRelays caps the auto pool (but not explicit relays).
417 func TestMOLSSelectPriorityMaxActiveRelaysLimitsAutoPool(t *testing.T) {
450 - policy := MOLSRelayPolicy{}
418
419 relays := make([]RelayState, 10)
420 for i := range relays {
421 relays[i] = confirmedPolicyRelayState(t, fmt.Sprintf("https://relay-%d.example", i))
422 }
423
457 - selected := policy.SelectPriority(relays, ClientState{MaxActiveRelays: 3})
424 + selected := selectPriority(relays, RouteState{MaxActiveRelays: 3})
425 if len(selected) != 3 {
426 t.Fatalf("len(selected) = %d, want 3", len(selected))
427 }
428 }
429
430 func TestMOLSSelectPriorityZeroMaxActiveRelaysUsesDefault(t *testing.T) {
464 - policy := MOLSRelayPolicy{}
431
432 relays := make([]RelayState, 10)
433 for i := range relays {
434 relays[i] = confirmedPolicyRelayState(t, fmt.Sprintf("https://relay-default-%d.example", i))
435 }
436
471 - selected := policy.SelectPriority(relays, ClientState{MaxActiveRelays: 0})
437 + selected := selectPriority(relays, RouteState{MaxActiveRelays: 0})
438 if len(selected) != defaultMaxActiveRelays {
439 t.Fatalf("len(selected) = %d, want %d", len(selected), defaultMaxActiveRelays)
440 }
441 }
442
443 func TestMOLSSelectPrioritySkipsExpiredAutoRelay(t *testing.T) {
478 - policy := MOLSRelayPolicy{}
444 expired := confirmedPolicyRelayState(t, "https://relay-expired.example")
445 expired.Descriptor.ExpiresAt = time.Now().UTC().Add(-time.Minute)
446
482 - if selected := policy.SelectPriority([]RelayState{expired}, ClientState{}); len(selected) != 0 {
447 + if selected := selectPriority([]RelayState{expired}, RouteState{}); len(selected) != 0 {
448 t.Fatalf("SelectPriority(expired auto) = %v, want empty", selected)
449 }
450 }
451
452 +func TestMOLSSelectPrioritySkipsBannedRelay(t *testing.T) {
453 + banned := confirmedPolicyRelayState(t, "https://relay-banned.example")
454 + banned.Banned = true
455 +
456 + if selected := selectPriority([]RelayState{banned}, RouteState{}); len(selected) != 0 {
457 + t.Fatalf("SelectPriority(banned) = %v, want empty", selected)
458 + }
459 +}
460 +
461 func TestMOLSSelectPriorityKeepsExpiredExplicitRelay(t *testing.T) {
488 - policy := MOLSRelayPolicy{}
462 relayURL := "https://relay-explicit-expired.example"
463 expired := confirmedPolicyRelayState(t, relayURL)
464 expired.Descriptor.ExpiresAt = time.Now().UTC().Add(-time.Minute)
465
493 - selected := policy.SelectPriority([]RelayState{expired}, ClientState{
466 + selected := selectPriority([]RelayState{expired}, RouteState{
467 ExplicitRelayURLs: []string{relayURL},
468 })
469 if len(selected) != 1 || selected[0] != relayURL {
@@ -499,39 +472,36 @@ func TestMOLSSelectPriorityKeepsExpiredExplicitRelay(t *testing.T) {
472 }
473
474 func TestMOLSSelectPrioritySkipsAutoRelayInBackoff(t *testing.T) {
502 - policy := MOLSRelayPolicy{}
475 backingOff := confirmedPolicyRelayState(t, "https://relay-backoff.example")
476 backingOff.suppressActiveUntil = time.Now().UTC().Add(time.Minute)
477
506 - if selected := policy.SelectPriority([]RelayState{backingOff}, ClientState{}); len(selected) != 0 {
478 + if selected := selectPriority([]RelayState{backingOff}, RouteState{}); len(selected) != 0 {
479 t.Fatalf("SelectPriority(backing off auto) = %v, want empty", selected)
480 }
481 }
482
483 func TestMOLSSelectPriorityKeepsDiscoveryBackoffRelay(t *testing.T) {
512 - policy := MOLSRelayPolicy{}
484 relayURL := "https://relay-discovery-backoff.example"
485 backingOff := confirmedPolicyRelayState(t, relayURL)
486 backingOff.nextDiscoveryRefreshAt = time.Now().UTC().Add(time.Minute)
487
517 - selected := policy.SelectPriority([]RelayState{backingOff}, ClientState{})
488 + selected := selectPriority([]RelayState{backingOff}, RouteState{})
489 if len(selected) != 1 || selected[0] != relayURL {
490 t.Fatalf("SelectPriority(discovery backoff) = %v, want [%q]", selected, relayURL)
491 }
492 }
493
494 func TestMOLSSelectPriorityKeepsUnobservedAutoSeed(t *testing.T) {
524 - policy := MOLSRelayPolicy{}
495 relayURL := "https://relay-seed.example"
496
527 - selected := policy.SelectPriority([]RelayState{bootstrapPolicyRelayState(relayURL)}, ClientState{})
497 + selected := selectPriority([]RelayState{bootstrapPolicyRelayState(relayURL)}, RouteState{})
498 if len(selected) != 1 || selected[0] != relayURL {
499 t.Fatalf("SelectPriority(unobserved seed) = %v, want [%q]", selected, relayURL)
500 }
501 }
502
503 // TestMOLSMagicRowSum verifies that each row of the base MOLS score grid sums
534 -// to the magic constant n*(n²+1)/2 = 131104.
504 +// to the magic constant n*(n^2+1)/2 = 131104.
505 func TestMOLSMagicRowSum(t *testing.T) {
506 const magicSum = molsOrder * (molsOrder*molsOrder + 1) / 2 // 131104
507
@@ -570,7 +540,7 @@ func TestMOLSMagicMainDiagonalSum(t *testing.T) {
540 for k := range uint8(64) {
541 diagSum += molsScore(k, k, molsBaseM1, molsBaseM2)
542 }
573 - // Allow ±1 rounding for floating-point-free integer arithmetic.
543 + // Allow +/-1 rounding for floating-point-free integer arithmetic.
544 diff := diagSum - magicSum
545 if diff < 0 {
546 diff = -diff
@@ -582,7 +552,7 @@ func TestMOLSMagicMainDiagonalSum(t *testing.T) {
552 }
553 }
554
585 -// TestMOLSGridUniqueness checks that all n² cells of the base grid have
555 +// TestMOLSGridUniqueness checks that all n^2 cells of the base grid have
556 // distinct values (Latin-square MOLS composite uniqueness).
557 func TestMOLSGridUniqueness(t *testing.T) {
558 seen := make(map[int]struct{}, 64*64)
@@ -619,7 +589,7 @@ func TestMOLSVariantGridUniqueness(t *testing.T) {
589
590 // TestMOLSHashToGF64InRange checks that hashToGF64 always returns [0, 63].
591 func TestMOLSHashToGF64InRange(t *testing.T) {
622 - inputs := []string{"", "a", "hello", "0x1234", "https://relay.example", "🔑"}
592 + inputs := []string{"", "a", "hello", "0x1234", "https://relay.example", "unicode-ish"}
593 for _, s := range inputs {
594 v := hashToGF64(s)
595 if v >= molsOrder {
portal/discovery/policy_test.go
+38 -41
@@ -1,7 +1,6 @@
1 package discovery
2
3 import (
4 - "errors"
4 "testing"
5 "time"
6
@@ -60,7 +59,6 @@ func confirmedPolicyRelayStateWithRTT(t *testing.T, relayURL string, rtt time.Du
59 }
60
61 func TestSelectPriorityMathematicalOrdering(t *testing.T) {
63 - policy := MOLSRelayPolicy{}
62 clientAddr := "192.168.0.10"
63 ingressIdx := hashToGF64(clientAddr)
64
@@ -75,7 +73,7 @@ func TestSelectPriorityMathematicalOrdering(t *testing.T) {
73 states = append(states, confirmedPolicyRelayState(t, url))
74 }
75
78 - selected := policy.SelectPriority(states, ClientState{LocalAddress: clientAddr})
76 + selected := selectPriority(states, RouteState{LocalAddress: clientAddr})
77
78 for i := 0; i < len(selected)-1; i++ {
79 scoreA := molsScore(ingressIdx, hashToGF64(selected[i]), molsBaseM1, molsBaseM2)
@@ -87,16 +85,15 @@ func TestSelectPriorityMathematicalOrdering(t *testing.T) {
85 }
86
87 func TestSelectPriorityKeepsExplicitRelaysOutsideAutoLimit(t *testing.T) {
90 - policy := MOLSRelayPolicy{}
88 explicitRelay := "https://relay-explicit.example"
89 relayA := "https://relay-a.example"
90 relayB := "https://relay-b.example"
91
95 - selected := policy.SelectPriority([]RelayState{
92 + selected := selectPriority([]RelayState{
93 bootstrapPolicyRelayState(explicitRelay),
94 confirmedPolicyRelayState(t, relayA),
95 confirmedPolicyRelayState(t, relayB),
99 - }, ClientState{
96 + }, RouteState{
97 LocalAddress: "127.0.0.1",
98 ExplicitRelayURLs: []string{explicitRelay},
99 MaxActiveRelays: 1,
@@ -111,7 +108,6 @@ func TestSelectPriorityKeepsExplicitRelaysOutsideAutoLimit(t *testing.T) {
108 }
109
110 func TestSelectPriorityCongestionInversion(t *testing.T) {
114 - policy := MOLSRelayPolicy{}
111 clientAddr := "10.0.0.1"
112 ingressIdx := hashToGF64(clientAddr)
113
@@ -121,7 +117,7 @@ func TestSelectPriorityCongestionInversion(t *testing.T) {
117 confirmedPolicyRelayStateWithRTT(t, r2, 800*time.Millisecond),
118 }
119
124 - selected := policy.SelectPriority(states, ClientState{LocalAddress: clientAddr})
120 + selected := selectPriority(states, RouteState{LocalAddress: clientAddr})
121
122 if len(selected) == 2 {
123 s1 := molsCongestionScore(ingressIdx, hashToGF64(selected[0]), molsBaseM1, molsBaseM2)
@@ -132,48 +128,39 @@ func TestSelectPriorityCongestionInversion(t *testing.T) {
128 }
129 }
130
135 -func TestSelectAggregateKeepsBootstrapRelayWhenDescriptorExpired(t *testing.T) {
136 - policy := MOLSRelayPolicy{}
137 - relayURL := "https://relay-bootstrap.example"
138 -
139 - state := bootstrapPolicyRelayState(relayURL)
140 - state.LastSeenAt = time.Now().UTC().Add(-time.Minute)
141 - state.Descriptor.ExpiresAt = time.Now().UTC().Add(-time.Second)
142 -
143 - selected := policy.SelectAggregate([]RelayState{state})
144 -
145 - if len(selected) != 1 {
146 - t.Fatalf("len(selected) = %d, want 1", len(selected))
147 - }
148 - if got := selected[0].Descriptor.APIHTTPSAddr; got != relayURL {
149 - t.Fatalf("selected[0] = %q, want %q", got, relayURL)
150 - }
151 -}
152 -
131 func TestSelectPriorityFallbackPromotion(t *testing.T) {
154 - policy := MOLSRelayPolicy{}
132 states := []RelayState{
133 confirmedPolicyRelayStateWithRTT(t, "https://f1.com", 3*time.Second),
134 confirmedPolicyRelayStateWithRTT(t, "https://f2.com", 4*time.Second),
135 }
136
160 - selected := policy.SelectPriority(states, ClientState{LocalAddress: "1.1.1.1"})
137 + selected := selectPriority(states, RouteState{LocalAddress: "1.1.1.1"})
138
139 if len(selected) < molsMinActiveNodes {
140 t.Errorf("Fallback promotion failed: got %d, want %d", len(selected), molsMinActiveNodes)
141 }
142 }
143
167 -func TestOnActiveConfirmedResetsActiveFailures(t *testing.T) {
168 - policy := MOLSRelayPolicy{}
144 +func TestConfirmRelayURLResetsActiveFailures(t *testing.T) {
145 + relayURL := "https://error.io"
146 + set := NewRelaySet(nil)
147 state := RelayState{
148 + Descriptor: types.RelayDescriptor{
149 + APIHTTPSAddr: relayURL,
150 + },
151 activeFailures: 5,
152 suppressActiveUntil: time.Now().UTC().Add(time.Minute),
153 Confirmed: false,
154 }
155 + set.mu.Lock()
156 + set.relays[relayURL] = state
157 + set.mu.Unlock()
158
175 - state = policy.OnActiveConfirmed(state)
159 + set.ConfirmRelayURL(relayURL)
160
161 + set.mu.RLock()
162 + state = set.relays[relayURL]
163 + set.mu.RUnlock()
164 if !state.Confirmed {
165 t.Fatal("Confirmed should be true")
166 }
@@ -185,20 +172,25 @@ func TestOnActiveConfirmedResetsActiveFailures(t *testing.T) {
172 }
173 }
174
188 -func TestOnDiscoveryFailureBackoff(t *testing.T) {
189 - policy := MOLSRelayPolicy{}
190 - state := confirmedPolicyRelayState(t, "https://error.io")
175 +func TestRecordDiscoveryFailureBackoff(t *testing.T) {
176 + relayURL := "https://error.io"
177 + set := NewRelaySet(nil)
178 + set.mu.Lock()
179 + set.relays[relayURL] = confirmedPolicyRelayState(t, relayURL)
180 + set.mu.Unlock()
181 budget := 3
182
183 start := time.Now()
184 for i := 0; i < budget; i++ {
195 - var backed bool
196 - state, backed, _ = policy.OnDiscoveryFailure(state, errors.New("err"), budget)
185 + backed, _, _ := set.RecordDiscoveryFailure(relayURL, budget)
186 if i < budget-1 && backed {
187 t.Fatal("Premature backoff")
188 }
189 }
190
191 + set.mu.RLock()
192 + state := set.relays[relayURL]
193 + set.mu.RUnlock()
194 if !state.nextDiscoveryRefreshAt.After(start) {
195 t.Fatal("discovery retry timer not scheduled")
196 }
@@ -207,16 +199,21 @@ func TestOnDiscoveryFailureBackoff(t *testing.T) {
199 }
200 }
201
210 -func TestOnActiveFailureBackoff(t *testing.T) {
211 - policy := MOLSRelayPolicy{}
212 - state := confirmedPolicyRelayState(t, "https://error.io")
202 +func TestRecordActiveFailureBackoff(t *testing.T) {
203 + relayURL := "https://error.io"
204 + set := NewRelaySet(nil)
205 + set.mu.Lock()
206 + set.relays[relayURL] = confirmedPolicyRelayState(t, relayURL)
207 + set.mu.Unlock()
208 start := time.Now()
209
215 - var backed bool
216 - state, backed, _ = policy.OnActiveFailure(state, errors.New("err"), 1)
210 + backed, _, _ := set.RecordActiveFailure(relayURL, 1)
211 if !backed {
212 t.Fatal("active failure should back off at budget")
213 }
214 + set.mu.RLock()
215 + state := set.relays[relayURL]
216 + set.mu.RUnlock()
217 if !state.suppressActiveUntil.After(start) {
218 t.Fatal("active suppression timer not scheduled")
219 }
portal/discovery/refresher.go
+1 -1
@@ -252,7 +252,7 @@ func (r *Refresher) refreshOverlay(ctx context.Context) error {
252 }
253
254 func (r *Refresher) logDiscoveryFailure(targetRelayURL, sourceURL string, recoveryFailures int, err error) {
255 - backedOff, backoffReason, failureCount := r.relaySet.RecordDiscoveryFailure(targetRelayURL, err, recoveryFailures)
255 + backedOff, backoffReason, failureCount := r.relaySet.RecordDiscoveryFailure(targetRelayURL, recoveryFailures)
256 if !backedOff {
257 return
258 }
portal/discovery/relayset.go
+94 -48
@@ -4,6 +4,7 @@ import (
4 "errors"
5 "fmt"
6 "reflect"
7 + "slices"
8 "sort"
9 "strings"
10 "sync"
@@ -46,7 +47,6 @@ type RelaySet struct {
47 mu sync.RWMutex
48 relays map[string]RelayState
49 keyIndex map[string]keyIndexEntry
49 - policy MOLSRelayPolicy
50 }
51
52 // keyIndexEntry records the rollback anchor for a signing identity.
@@ -71,7 +71,6 @@ func NewRelaySet(bootstrapRelayURLs []string) *RelaySet {
71 set := &RelaySet{
72 relays: make(map[string]RelayState),
73 keyIndex: make(map[string]keyIndexEntry),
74 - policy: MOLSRelayPolicy{},
74 }
75 set.SetBootstrapRelayURLs(bootstrapRelayURLs)
76 return set
@@ -154,15 +153,6 @@ func (s *RelaySet) upsertDescriptorLocked(record RelayState, now time.Time, allo
153 return upsertAccepted
154 }
155
157 -func (s *RelaySet) SetRelayPolicy(policy MOLSRelayPolicy) {
158 - if policy == (MOLSRelayPolicy{}) {
159 - policy = MOLSRelayPolicy{}
160 - }
161 - s.mu.Lock()
162 - defer s.mu.Unlock()
163 - s.policy = policy
164 -}
165 -
156 func (s *RelaySet) SetBootstrapRelayURLs(inputs []string) {
157 s.mu.Lock()
158 defer s.mu.Unlock()
@@ -196,52 +186,78 @@ func (s *RelaySet) SetBootstrapRelayURLs(inputs []string) {
186 }
187 }
188
199 -func (s *RelaySet) AggregateRelays() []RelayState {
200 - s.mu.RLock()
201 - states := make([]RelayState, 0, len(s.relays))
202 - for _, state := range s.relays {
203 - states = append(states, state)
189 +type Route struct {
190 + path []string
191 + explicit bool
192 +}
193 +
194 +func NewRoute(path []string, explicit bool) Route {
195 + return Route{
196 + path: append([]string(nil), path...),
197 + explicit: explicit,
198 }
205 - policy := s.policy
206 - s.mu.RUnlock()
199 +}
200
208 - return policy.SelectAggregate(states)
201 +func (r Route) Explicit() bool {
202 + return r.explicit
203 }
204
211 -func (s *RelaySet) ConfirmedRelays() []RelayState {
212 - s.mu.RLock()
213 - states := make([]RelayState, 0, len(s.relays))
214 - for _, state := range s.relays {
215 - states = append(states, state)
205 +func (r Route) ListenerRelayURL() string {
206 + if len(r.path) == 0 {
207 + return ""
208 }
217 - policy := s.policy
218 - s.mu.RUnlock()
219 -
220 - return policy.SelectConfirmed(states)
209 + return r.path[len(r.path)-1]
210 }
211
223 -func (s *RelaySet) PriorityRelays(clientState ClientState) []string {
224 - s.mu.RLock()
225 - states := make([]RelayState, 0, len(s.relays))
226 - for _, state := range s.relays {
227 - states = append(states, state)
212 +func (r Route) MultiHop() []string {
213 + if len(r.path) <= 1 {
214 + return nil
215 }
229 - policy := s.policy
230 - s.mu.RUnlock()
216 + return append([]string(nil), r.path...)
217 +}
218
232 - return policy.SelectPriority(states, clientState)
219 +func (r Route) Equal(other Route) bool {
220 + return r.explicit == other.explicit && slices.Equal(r.path, other.path)
221 }
222
235 -func (s *RelaySet) PriorityMultiHop(clientState ClientState) []string {
223 +func (r Route) WithListenerRelayURL(relayURL string) Route {
224 + if len(r.path) == 0 {
225 + return NewRoute([]string{relayURL}, r.explicit)
226 + }
227 + path := append([]string(nil), r.path...)
228 + path[len(path)-1] = relayURL
229 + return NewRoute(path, r.explicit)
230 +}
231 +
232 +func (s *RelaySet) PlanRoutes(explicitPath []string, routeState RouteState) ([]Route, error) {
233 + if len(explicitPath) > 0 {
234 + if len(explicitPath) == 1 {
235 + return nil, fmt.Errorf("multi-hop requires at least entry and exit relay urls")
236 + }
237 + return []Route{NewRoute(explicitPath, true)}, nil
238 + }
239 +
240 s.mu.RLock()
241 states := make([]RelayState, 0, len(s.relays))
242 for _, state := range s.relays {
243 states = append(states, state)
244 }
241 - policy := s.policy
245 s.mu.RUnlock()
246
244 - return policy.SelectMultiHop(states, clientState)
247 + if routeState.MultiHopDepth > 1 {
248 + path := selectMultiHop(states, routeState)
249 + if len(path) < routeState.MultiHopDepth {
250 + return nil, fmt.Errorf("multi-hop-depth %d requires %d overlay relay candidates, got %d", routeState.MultiHopDepth, routeState.MultiHopDepth, len(path))
251 + }
252 + return []Route{NewRoute(path, false)}, nil
253 + }
254 +
255 + relayURLs := selectPriority(states, routeState)
256 + routes := make([]Route, 0, len(relayURLs))
257 + for _, relayURL := range relayURLs {
258 + routes = append(routes, NewRoute([]string{relayURL}, slices.Contains(routeState.ExplicitRelayURLs, relayURL)))
259 + }
260 + return routes, nil
261 }
262
263 func (s *RelaySet) OverlayPeerStates() []RelayState {
@@ -344,7 +360,7 @@ func (s *RelaySet) BanRelayURL(relayURL string) {
360 if !ok {
361 state = newRelayState(relayURL)
362 }
347 - state = s.policy.OnBanned(state)
363 + state.Banned = true
364 s.relays[relayURL] = state
365 }
366
@@ -356,7 +372,9 @@ func (s *RelaySet) ConfirmRelayURL(relayURL string) {
372 if !ok {
373 state = newRelayState(relayURL)
374 }
359 - state = s.policy.OnActiveConfirmed(state)
375 + state.Confirmed = true
376 + state.activeFailures = 0
377 + state.suppressActiveUntil = time.Time{}
378 s.relays[relayURL] = state
379 }
380
@@ -368,7 +386,7 @@ func (s *RelaySet) UnconfirmRelayURL(relayURL string) {
386 if !ok {
387 return
388 }
371 - state = s.policy.OnUnconfirmed(state)
389 + state.Confirmed = false
390 s.relays[relayURL] = state
391 }
392
@@ -442,7 +460,8 @@ func (s *RelaySet) ApplyRelayDiscoveryResponse(targetURL string, resp types.Disc
460
461 isAuthoritativeTarget := !protocolMismatch && !missingTarget && authoritative && relayURL == targetURL
462 if isAuthoritativeTarget {
445 - record = s.policy.OnDiscoveryConfirmed(record)
463 + record.discoveryFailures = 0
464 + record.nextDiscoveryRefreshAt = time.Time{}
465 }
466
467 if upsert := s.upsertDescriptorLocked(record, now, isAuthoritativeTarget); upsert != upsertAccepted {
@@ -454,7 +473,8 @@ func (s *RelaySet) ApplyRelayDiscoveryResponse(targetURL string, resp types.Disc
473 // existing URL slot.
474 if isAuthoritativeTarget && hasExistingAtURL {
475 if existingAtURL.discoveryFailures != 0 || !existingAtURL.nextDiscoveryRefreshAt.IsZero() {
457 - existingAtURL = s.policy.OnDiscoveryConfirmed(existingAtURL)
476 + existingAtURL.discoveryFailures = 0
477 + existingAtURL.nextDiscoveryRefreshAt = time.Time{}
478 s.relays[relayURL] = existingAtURL
479 relaySetChanged = true
480 }
@@ -634,7 +654,7 @@ func (s *RelaySet) enforceCapLocked() {
654 }
655 }
656
637 -func (s *RelaySet) RecordDiscoveryFailure(relayURL string, err error, recoveryFailures int) (backedOff bool, backoffReason string, failureCount int) {
657 +func (s *RelaySet) RecordDiscoveryFailure(relayURL string, recoveryFailures int) (backedOff bool, backoffReason string, failureCount int) {
658 s.mu.Lock()
659 defer s.mu.Unlock()
660
@@ -642,12 +662,25 @@ func (s *RelaySet) RecordDiscoveryFailure(relayURL string, err error, recoveryFa
662 if !ok {
663 return false, "", 0
664 }
645 - state, backedOff, backoffReason = s.policy.OnDiscoveryFailure(state, err, recoveryFailures)
665 + state.discoveryFailures++
666 +
667 + if recoveryFailures <= 0 || state.discoveryFailures < recoveryFailures {
668 + s.relays[relayURL] = state
669 + return false, "retry", state.discoveryFailures
670 + }
671 + failuresOverBudget := state.discoveryFailures - recoveryFailures
672 + backoff := defaultDirectRecoveryBackoff << min(failuresOverBudget, 3)
673 + if backoff > maxDirectRecoveryBackoff {
674 + backoff = maxDirectRecoveryBackoff
675 + }
676 + state.nextDiscoveryRefreshAt = time.Now().Add(backoff)
677 + backedOff = true
678 + backoffReason = "discovery"
679 s.relays[relayURL] = state
680 return backedOff, backoffReason, state.discoveryFailures
681 }
682
650 -func (s *RelaySet) RecordActiveFailure(relayURL string, err error, recoveryFailures int) (backedOff bool, backoffReason string, failureCount int) {
683 +func (s *RelaySet) RecordActiveFailure(relayURL string, recoveryFailures int) (backedOff bool, backoffReason string, failureCount int) {
684 s.mu.Lock()
685 defer s.mu.Unlock()
686
@@ -655,7 +688,20 @@ func (s *RelaySet) RecordActiveFailure(relayURL string, err error, recoveryFailu
688 if !ok {
689 return false, "", 0
690 }
658 - state, backedOff, backoffReason = s.policy.OnActiveFailure(state, err, recoveryFailures)
691 + state.activeFailures++
692 +
693 + if recoveryFailures <= 0 || state.activeFailures < recoveryFailures {
694 + s.relays[relayURL] = state
695 + return false, "retry", state.activeFailures
696 + }
697 + failuresOverBudget := state.activeFailures - recoveryFailures
698 + backoff := defaultDirectRecoveryBackoff << min(failuresOverBudget, 3)
699 + if backoff > maxDirectRecoveryBackoff {
700 + backoff = maxDirectRecoveryBackoff
701 + }
702 + state.suppressActiveUntil = time.Now().Add(backoff)
703 + backedOff = true
704 + backoffReason = "active"
705 s.relays[relayURL] = state
706 return backedOff, backoffReason, state.activeFailures
707 }
portal/discovery/relayset_test.go
+45 -6
@@ -7,6 +7,18 @@ import (
7 "github.com/gosuda/portal-tunnel/v2/types"
8 )
9
10 +func relayStates(set *RelaySet) []RelayState {
11 + set.mu.RLock()
12 + defer set.mu.RUnlock()
13 + states := make([]RelayState, 0, len(set.relays))
14 + for _, state := range set.relays {
15 + if !state.Banned {
16 + states = append(states, state)
17 + }
18 + }
19 + return states
20 +}
21 +
22 func TestApplyRelayDiscoveryResponsePreservesBootstrapFlag(t *testing.T) {
23 set := NewRelaySet([]string{"https://relay-a.example"})
24
@@ -18,9 +30,9 @@ func TestApplyRelayDiscoveryResponsePreservesBootstrapFlag(t *testing.T) {
30 t.Fatalf("ApplyRelayDiscoveryResponse() error = %v", err)
31 }
32
21 - states := set.AggregateRelays()
33 + states := relayStates(set)
34 if len(states) != 1 {
23 - t.Fatalf("len(AggregateRelays()) = %d, want 1", len(states))
35 + t.Fatalf("len(relayStates()) = %d, want 1", len(states))
36 }
37 if !states[0].Bootstrap {
38 t.Fatal("bootstrap relay lost bootstrap flag after discovery update")
@@ -63,9 +75,9 @@ func TestApplyRelayDiscoveryResponseCollectsRelaysDespiteProtocolMismatch(t *tes
75 t.Fatal("expected protocol-mismatched discovery response to change relay set")
76 }
77
66 - states := set.AggregateRelays()
78 + states := relayStates(set)
79 if len(states) != 1 {
68 - t.Fatalf("len(AggregateRelays()) = %d, want 1", len(states))
80 + t.Fatalf("len(relayStates()) = %d, want 1", len(states))
81 }
82 if got := states[0].Descriptor.APIHTTPSAddr; got != desc.APIHTTPSAddr {
83 t.Fatalf("states[0] = %q, want %q", got, desc.APIHTTPSAddr)
@@ -90,9 +102,9 @@ func TestApplyRelayDiscoveryResponseCollectsHintsWhenTargetDescriptorIsMissing(t
102 t.Fatal("expected hinted relay to still be collected")
103 }
104
93 - states := set.AggregateRelays()
105 + states := relayStates(set)
106 if len(states) != 1 {
95 - t.Fatalf("len(AggregateRelays()) = %d, want 1", len(states))
107 + t.Fatalf("len(relayStates()) = %d, want 1", len(states))
108 }
109 if got := states[0].Descriptor.APIHTTPSAddr; got != hinted.APIHTTPSAddr {
110 t.Fatalf("states[0] = %q, want %q", got, hinted.APIHTTPSAddr)
@@ -218,3 +230,30 @@ func TestUnconfirmRelayURLClearsLocalConfirmationOnly(t *testing.T) {
230 t.Fatal("relay should lose local confirmation after listener failure")
231 }
232 }
233 +
234 +func TestPlanRoutesExplicitPathReturnsSingleRouteToExit(t *testing.T) {
235 + const (
236 + entry = "https://entry.example"
237 + mid = "https://middle.example"
238 + exit = "https://exit.example"
239 + )
240 +
241 + routes, err := NewRelaySet(nil).PlanRoutes([]string{entry, mid, exit}, RouteState{})
242 + if err != nil {
243 + t.Fatalf("PlanRoutes() error = %v", err)
244 + }
245 + if len(routes) != 1 {
246 + t.Fatalf("len(routes) = %d, want 1", len(routes))
247 + }
248 + route := routes[0]
249 + if !route.Explicit() {
250 + t.Fatal("route.Explicit() = false, want true")
251 + }
252 + if got := route.ListenerRelayURL(); got != exit {
253 + t.Fatalf("ListenerRelayURL() = %q, want %q", got, exit)
254 + }
255 + path := route.MultiHop()
256 + if len(path) != 3 || path[0] != entry || path[1] != mid || path[2] != exit {
257 + t.Fatalf("MultiHop() = %v, want [%q %q %q]", path, entry, mid, exit)
258 + }
259 +}
portal/discovery/relaystate.go
+2 -2
@@ -57,7 +57,7 @@ func (state RelayState) hasObservedDescriptor() bool {
57 return !state.LastSeenAt.IsZero()
58 }
59
60 -type ClientState struct {
60 +type RouteState struct {
61 ExplicitRelayURLs []string
62 // MaxActiveRelays caps auto-selected relays. Zero or negative values use
63 // the policy default of 3.
@@ -65,7 +65,7 @@ type ClientState struct {
65 MultiHopDepth int
66 RequireUDP bool
67 RequireTCP bool
68 - // LocalAddress is the ingress identity address used by MOLSRelayPolicy to
68 + // LocalAddress is the ingress identity address used by MOLS route selection to
69 // derive a deterministic row index into the GF(64) MOLS grid.
70 LocalAddress string
71 }
sdk/api_client.go
+9 -7
@@ -90,8 +90,9 @@ func (l *listener) registerLease(ctx context.Context, ttl time.Duration, udpEnab
90 var publicHostname string
91 var keylessURL string
92 var hopRoutes []types.HopRoute
93 - if len(l.multiHop) > 0 {
94 - if len(l.multiHop) < 2 {
93 + multiHop := l.route.MultiHop()
94 + if len(multiHop) > 0 {
95 + if len(multiHop) < 2 {
96 return types.RegisterResponse{}, nil, errors.New("multi-hop requires at least entry and exit relay urls")
97 }
98 if l.relaySet == nil {
@@ -99,8 +100,8 @@ func (l *listener) registerLease(ctx context.Context, ttl time.Duration, udpEnab
100 }
101
102 now := time.Now().UTC()
102 - hopPath := make([]types.RelayDescriptor, 0, len(l.multiHop))
103 - for i, relayURL := range l.multiHop {
103 + hopPath := make([]types.RelayDescriptor, 0, len(multiHop))
104 + for i, relayURL := range multiHop {
105 desc, ok := l.relaySet.OverlayRelayDescriptor(relayURL, now)
106 if !ok {
107 return types.RegisterResponse{}, nil, fmt.Errorf("multi-hop relay %d descriptor is unavailable", i)
@@ -108,12 +109,13 @@ func (l *listener) registerLease(ctx context.Context, ttl time.Duration, udpEnab
109 hopPath = append(hopPath, desc)
110 }
111
112 + entryRelayURL := hopPath[0].APIHTTPSAddr
113 var err error
112 - publicHostname, err = utils.LeaseHostname(l.identity.Name, utils.PortalRootHost(hopPath[0].APIHTTPSAddr))
114 + publicHostname, err = utils.LeaseHostname(l.identity.Name, utils.PortalRootHost(entryRelayURL))
115 if err != nil {
116 return types.RegisterResponse{}, nil, err
117 }
116 - keylessURL = hopPath[0].APIHTTPSAddr
118 + keylessURL = entryRelayURL
119
120 hopRoutes = make([]types.HopRoute, 0, len(hopPath)-1)
121 var previousHopToken string
@@ -136,7 +138,7 @@ func (l *listener) registerLease(ctx context.Context, ttl time.Duration, udpEnab
138 }
139 if i == 0 {
140 route.MatchHostname = publicHostname
139 - route.Metadata = l.metadata
141 + route.Metadata = l.metadata.Copy()
142 } else {
143 route.MatchToken = previousHopToken
144 }
sdk/expose.go
+56 -55
@@ -101,12 +101,13 @@ func Expose(ctx context.Context, cfg ExposeConfig) (*Exposure, error) {
101 return nil, errors.New("multi-hop currently supports only the default SNI TLS stream transport")
102 }
103
104 - var listenerRelayURLs []string
104 + var initialRouteCount int
105 var relaySetURLs []string
106 if len(multiHop) > 0 {
107 - listenerRelayURLs = []string{multiHop[len(multiHop)-1]}
107 + initialRouteCount = 1
108 relaySetURLs = append([]string(nil), multiHop...)
109 } else if cfg.MultiHopDepth > 1 {
110 + initialRouteCount = 1
111 relaySetURLs, err = utils.ResolvePortalRelayURLs(ctx, explicitRelayURLs, cfg.Discovery)
112 if err != nil {
113 return nil, err
@@ -116,7 +117,7 @@ func Expose(ctx context.Context, cfg ExposeConfig) (*Exposure, error) {
117 if err != nil {
118 return nil, err
119 }
119 - listenerRelayURLs = append([]string(nil), explicitRelayURLs...)
120 + initialRouteCount = len(explicitRelayURLs)
121 }
122
123 identity, createdIdentity, err := utils.ResolveListenerIdentity(
@@ -160,10 +161,10 @@ func Expose(ctx context.Context, cfg ExposeConfig) (*Exposure, error) {
161 banMITM: cfg.BanMITM,
162 maxActiveRelays: cfg.MaxActiveRelays,
163 metadata: cfg.Metadata,
163 - accepted: make(chan net.Conn, max(initialRouteCapacity(listenerRelayURLs, cfg.MultiHopDepth)*defaultReadyTarget*2, 1)),
164 - datagrams: make(chan types.DatagramFrame, max(initialRouteCapacity(listenerRelayURLs, cfg.MultiHopDepth)*32, 1)),
164 + accepted: make(chan net.Conn, max(initialRouteCount*defaultReadyTarget*2, 1)),
165 + datagrams: make(chan types.DatagramFrame, max(initialRouteCount*32, 1)),
166 relaySet: discovery.NewRelaySet(relaySetURLs),
166 - relayListeners: make(map[string]*listener, initialRouteCapacity(listenerRelayURLs, cfg.MultiHopDepth)),
167 + relayListeners: make(map[string]*listener, initialRouteCount),
168 }
169
170 if cfg.Discovery || len(multiHop) > 0 || cfg.MultiHopDepth > 1 {
@@ -174,7 +175,7 @@ func Expose(ctx context.Context, cfg ExposeConfig) (*Exposure, error) {
175 }
176 }
177
177 - if len(listenerRelayURLs) > 0 || cfg.Discovery || cfg.MultiHopDepth > 1 {
178 + if initialRouteCount > 0 || cfg.Discovery {
179 if err := exposure.reconcileRelayListeners(true); err != nil {
180 _ = exposure.Close()
181 return nil, err
@@ -193,13 +194,6 @@ func Expose(ctx context.Context, cfg ExposeConfig) (*Exposure, error) {
194 return exposure, nil
195 }
196
196 -func initialRouteCapacity(listenerRelayURLs []string, multiHopDepth int) int {
197 - if multiHopDepth > 1 {
198 - return 1
199 - }
200 - return len(listenerRelayURLs)
201 -}
202 -
197 func (e *Exposure) ActiveRelayURLs() []string {
198 e.listenerMu.RLock()
199 defer e.listenerMu.RUnlock()
@@ -438,56 +432,52 @@ func (e *Exposure) runDiscoveryLoop(ctx context.Context) {
432 }
433
434 func (e *Exposure) reconcileRelayListeners(failOnError bool) error {
441 - var listenerRelayURLs []string
442 - var multiHop []string
443 - if len(e.multiHop) > 0 {
444 - listenerRelayURLs = []string{e.multiHop[len(e.multiHop)-1]}
445 - multiHop = append([]string(nil), e.multiHop...)
446 - } else if e.multiHopDepth > 1 {
447 - multiHop = e.relaySet.PriorityMultiHop(discovery.ClientState{
448 - MultiHopDepth: e.multiHopDepth,
449 - LocalAddress: e.identity.Address,
450 - })
451 - if len(multiHop) < e.multiHopDepth {
452 - return fmt.Errorf("multi-hop-depth %d requires %d overlay relay candidates, got %d", e.multiHopDepth, e.multiHopDepth, len(multiHop))
435 + if e.relaySet == nil {
436 + return errors.New("relay set is unavailable")
437 + }
438 + routes, err := e.relaySet.PlanRoutes(append([]string(nil), e.multiHop...), discovery.RouteState{
439 + ExplicitRelayURLs: append([]string(nil), e.explicitRelays...),
440 + MaxActiveRelays: e.maxActiveRelays,
441 + MultiHopDepth: e.multiHopDepth,
442 + RequireUDP: e.udpEnabled,
443 + RequireTCP: e.tcpEnabled,
444 + LocalAddress: e.identity.Address,
445 + })
446 + if err != nil {
447 + return err
448 + }
449 +
450 + routesByRelay := make(map[string]discovery.Route, len(routes))
451 + for _, route := range routes {
452 + relayURL := route.ListenerRelayURL()
453 + if relayURL == "" {
454 + continue
455 }
454 - listenerRelayURLs = []string{multiHop[len(multiHop)-1]}
455 - } else {
456 - listenerRelayURLs = e.relaySet.PriorityRelays(discovery.ClientState{
457 - ExplicitRelayURLs: append([]string(nil), e.explicitRelays...),
458 - MaxActiveRelays: e.maxActiveRelays,
459 - RequireUDP: e.udpEnabled,
460 - RequireTCP: e.tcpEnabled,
461 - LocalAddress: e.identity.Address,
462 - })
456 + routesByRelay[relayURL] = route
457 }
458
459 e.listenerMu.Lock()
466 - staleRelayListeners := make(map[string]*listener)
467 - removedRelayURLs := make([]string, 0)
460 + staleListeners := make(map[string]*listener)
461 for relayURL, listener := range e.relayListeners {
469 - if slices.Contains(listenerRelayURLs, relayURL) && slices.Equal(listener.multiHop, multiHop) {
462 + route, wanted := routesByRelay[relayURL]
463 + if wanted && listener != nil && listener.route.Equal(route) {
464 continue
465 }
472 - staleRelayListeners[relayURL] = listener
473 - removedRelayURLs = append(removedRelayURLs, relayURL)
466 + staleListeners[relayURL] = listener
467 delete(e.relayListeners, relayURL)
468 }
476 -
477 - missingRelayURLs := make([]string, 0, len(listenerRelayURLs))
478 - for _, relayURL := range listenerRelayURLs {
479 - if _, ok := e.relayListeners[relayURL]; ok {
469 + missingRoutes := make([]discovery.Route, 0)
470 + for _, route := range routes {
471 + relayURL := route.ListenerRelayURL()
472 + if _, exists := e.relayListeners[relayURL]; exists {
473 continue
474 }
482 - missingRelayURLs = append(missingRelayURLs, relayURL)
475 + missingRoutes = append(missingRoutes, route)
476 }
477 e.listenerMu.Unlock()
485 - if len(removedRelayURLs) > 1 {
486 - slices.Sort(removedRelayURLs)
487 - }
478
489 - addedRelayURLs := make([]string, 0, len(missingRelayURLs))
490 - for relayURL, listener := range staleRelayListeners {
479 + addedRelayURLs := make([]string, 0, len(missingRoutes))
480 + for relayURL, listener := range staleListeners {
481 if listener == nil {
482 continue
483 }
@@ -495,16 +485,16 @@ func (e *Exposure) reconcileRelayListeners(failOnError bool) error {
485 log.Warn().Err(err).Str("relay_url", relayURL).Msg("close stale relay listener")
486 }
487 }
498 - for _, relayURL := range missingRelayURLs {
488 + for _, route := range missingRoutes {
489 + relayURL := route.ListenerRelayURL()
490 retryCount := 10
500 - if len(multiHop) > 0 || slices.Contains(e.explicitRelays, relayURL) {
491 + if route.Explicit() || len(route.MultiHop()) > 0 {
492 retryCount = 0
493 }
503 - listener, err := newListener(context.Background(), relayURL, listenerConfig{
494 + listener, err := newListener(context.Background(), route, listenerConfig{
495 Identity: e.identity,
496 UDPEnabled: e.udpEnabled,
497 TCPEnabled: e.tcpEnabled,
507 - MultiHop: multiHop,
498 BanMITM: e.banMITM,
499 RetryCount: retryCount,
500 Metadata: e.metadata,
@@ -538,7 +528,18 @@ func (e *Exposure) reconcileRelayListeners(failOnError bool) error {
528 go e.runListenerAcceptLoop(listener)
529 }
530
541 - if len(removedRelayURLs) > 0 || len(addedRelayURLs) > 0 {
531 + if len(staleListeners) > 0 || len(addedRelayURLs) > 0 {
532 + removedRelayURLs := make([]string, 0, len(staleListeners))
533 + for relayURL := range staleListeners {
534 + removedRelayURLs = append(removedRelayURLs, relayURL)
535 + }
536 + if len(removedRelayURLs) > 1 {
537 + slices.Sort(removedRelayURLs)
538 + }
539 + listenerRelayURLs := make([]string, 0, len(routes))
540 + for _, route := range routes {
541 + listenerRelayURLs = append(listenerRelayURLs, route.ListenerRelayURL())
542 + }
543 log.Info().
544 Strs("added_relays", addedRelayURLs).
545 Strs("removed_relays", removedRelayURLs).
sdk/expose_test.go
+4
@@ -36,11 +36,13 @@ func TestExposureReconcileRemovesBannedRelayFromActiveSet(t *testing.T) {
36 exposure.relayListeners = map[string]*listener{
37 relayA: {
38 relayURL: relayURL,
39 + route: discovery.NewRoute([]string{relayA}, true),
40 cancel: func() { close(relayAClosed) },
41 doneCh: relayAClosed,
42 },
43 relayB: {
44 relayURL: relayBURL,
45 + route: discovery.NewRoute([]string{relayB}, true),
46 },
47 }
48
@@ -91,11 +93,13 @@ func TestExposureReconcileRemovesStaleListener(t *testing.T) {
93 exposure.relayListeners = map[string]*listener{
94 relayA: {
95 relayURL: relayAURL,
96 + route: discovery.NewRoute([]string{relayA}, true),
97 cancel: func() { close(relayAClosed) },
98 doneCh: relayAClosed,
99 },
100 relayB: {
101 relayURL: relayBURL,
102 + route: discovery.NewRoute([]string{relayB}, true),
103 },
104 }
105
sdk/listener.go
+6 -7
@@ -40,7 +40,6 @@ type listenerConfig struct {
40 ReadyTarget int
41 RetryCount int
42 RetryWait time.Duration
43 - MultiHop []string
43 relaySet *discovery.RelaySet
44 }
45
@@ -52,10 +51,10 @@ type listener struct {
51 closeOnce sync.Once
52
53 relayURL *url.URL
54 + route discovery.Route
55 identity types.Identity
56 metadata types.LeaseMetadata
57 relaySet *discovery.RelaySet
58 - multiHop []string
58 udpEnabled bool
59 tcpEnabled bool
60 dialTimeout time.Duration
@@ -79,7 +78,7 @@ type listener struct {
78
79 // newListener creates one relay listener and its dedicated relay transport for one relay URL.
80 // Only local config validation fails immediately; relay startup runs in the background until ready.
82 -func newListener(ctx context.Context, relayURL string, cfg listenerConfig) (*listener, error) {
81 +func newListener(ctx context.Context, route discovery.Route, cfg listenerConfig) (*listener, error) {
82 listenerCtx, cancel := context.WithCancel(ctx)
83 readyTarget := utils.IntOrDefault(cfg.ReadyTarget, defaultReadyTarget)
84 leaseTTL := utils.DurationOrDefault(cfg.LeaseTTL, defaultLeaseTTL)
@@ -89,7 +88,7 @@ func newListener(ctx context.Context, relayURL string, cfg listenerConfig) (*lis
88 renewBefore := utils.DurationOrDefault(cfg.RenewBefore, defaultRenewBefore)
89 retryWait := utils.DurationOrDefault(cfg.RetryWait, defaultRetryWait)
90
92 - normalizedRelayURL, err := utils.NormalizeRelayURL(relayURL)
91 + normalizedRelayURL, err := utils.NormalizeRelayURL(route.ListenerRelayURL())
92 if err != nil {
93 cancel()
94 return nil, err
@@ -103,10 +102,10 @@ func newListener(ctx context.Context, relayURL string, cfg listenerConfig) (*lis
102 cancel: cancel,
103 doneCh: listenerCtx.Done(),
104 relayURL: relayurl,
105 + route: route.WithListenerRelayURL(normalizedRelayURL),
106 identity: cfg.Identity,
107 metadata: cfg.Metadata,
108 relaySet: cfg.relaySet,
109 - multiHop: cfg.MultiHop,
109 udpEnabled: cfg.UDPEnabled,
110 tcpEnabled: cfg.TCPEnabled,
111 dialTimeout: dialTimeout,
@@ -151,7 +150,7 @@ func (l *listener) run(ctx context.Context) {
150 relayURL := l.relayURL.String()
151 if l.relaySet != nil && relayURL != "" {
152 l.relaySet.UnconfirmRelayURL(relayURL)
154 - l.relaySet.RecordActiveFailure(relayURL, err, 1)
153 + l.relaySet.RecordActiveFailure(relayURL, 1)
154 }
155 log.Error().
156 Err(err).
@@ -839,7 +838,7 @@ func (l *listener) waitRetry(ctx context.Context, operation string, err error, r
838 if l.retryCount > 0 && retries > l.retryCount {
839 if l.relaySet != nil && relayURL != "" {
840 l.relaySet.UnconfirmRelayURL(relayURL)
842 - l.relaySet.RecordActiveFailure(relayURL, err, 1)
841 + l.relaySet.RecordActiveFailure(relayURL, 1)
842 }
843 logger.Error().
844 Err(err).
sdk/mitm_test.go
+14 -5
@@ -18,6 +18,7 @@ import (
18 "testing"
19 "time"
20
21 + "github.com/gosuda/portal-tunnel/v2/portal/discovery"
22 "github.com/gosuda/portal-tunnel/v2/types"
23 )
24
@@ -223,8 +224,12 @@ func TestMITMProbeDetectionBansListener(t *testing.T) {
224 Reason: types.MITMProbeReasonExporterMismatch,
225 }, nil)
226
226 - for _, state := range listener.relaySet.AggregateRelays() {
227 - if state.Descriptor.APIHTTPSAddr == relayURL.String() {
227 + routes, err := listener.relaySet.PlanRoutes(nil, discovery.RouteState{})
228 + if err != nil {
229 + t.Fatalf("PlanRoutes() error = %v", err)
230 + }
231 + for _, route := range routes {
232 + if route.ListenerRelayURL() == relayURL.String() {
233 t.Fatal("relay still active after mitm detection")
234 }
235 }
@@ -255,9 +260,13 @@ func TestMITMProbeDetectionWarnsWithoutBanningListener(t *testing.T) {
260 Reason: types.MITMProbeReasonExporterMismatch,
261 }, nil)
262
258 - activeRelayURLs := make([]string, 0)
259 - for _, state := range listener.relaySet.AggregateRelays() {
260 - activeRelayURLs = append(activeRelayURLs, state.Descriptor.APIHTTPSAddr)
263 + routes, err := listener.relaySet.PlanRoutes(nil, discovery.RouteState{})
264 + if err != nil {
265 + t.Fatalf("PlanRoutes() error = %v", err)
266 + }
267 + activeRelayURLs := make([]string, 0, len(routes))
268 + for _, route := range routes {
269 + activeRelayURLs = append(activeRelayURLs, route.ListenerRelayURL())
270 }
271 if len(activeRelayURLs) != 1 || activeRelayURLs[0] != relayURL.String() {
272 t.Fatalf("ActiveRelayURLs() = %v, want [%q]", activeRelayURLs, relayURL.String())