feat: enhance relay discovery and selection logic, add support for discovery flag in relay descriptors
Kim committed
Apr 13, 2026 at 17:04 UTC
7a0608be29c6b095f86a5dbfe41d058f12b55a72
12 files changed
+253
-323
portal/api_server.go
+2
-15
@@ -152,10 +152,6 @@ func (s *Server) extractAllowedClientIP(w http.ResponseWriter, r *http.Request)
152
}
153
154
func (s *Server) handleRelayDiscovery(w http.ResponseWriter, r *http.Request) {
155
- if !utils.RequireMethod(w, r, http.MethodGet) {
156
- return
157
- }
158
-
155
now := time.Now().UTC()
156
activeConns := float64(s.proxy.ActiveConns())
157
tcpTrafficBPS := s.proxy.CurrentTCPBPS(now)
@@ -182,6 +178,7 @@ func (s *Server) handleRelayDiscovery(w http.ResponseWriter, r *http.Request) {
178
IssuedAt: now,
179
ExpiresAt: now.Add(2 * discovery.DiscoveryPollInterval),
180
APIHTTPSAddr: s.cfg.PortalURL,
181
+ Discovery: s.cfg.DiscoveryEnabled,
182
IngressTLSAddr: ingressAddr,
183
WireGuardPublicKey: wireGuardPublicKey,
184
WireGuardEndpoint: wireGuardEndpoint,
@@ -198,17 +195,7 @@ func (s *Server) handleRelayDiscovery(w http.ResponseWriter, r *http.Request) {
195
utils.WriteAPIError(w, http.StatusInternalServerError, types.APIErrorCodeInternal, err.Error())
196
return
197
}
201
-
202
- resp := types.DiscoveryResponse{
203
- ProtocolVersion: types.ProtocolVersion,
204
- GeneratedAt: now,
205
- Self: self,
206
- Relays: nil,
207
- }
208
- if s.relaySet != nil {
209
- resp.Relays = s.relaySet.ConfirmedDescriptors()
210
- }
211
- utils.WriteAPIData(w, http.StatusOK, resp)
198
+ s.relaySet.ServeDiscovery(w, r, self)
199
}
200
201
func (s *Server) handleDomain(w http.ResponseWriter, r *http.Request) {
portal/discovery/policy.go
+2
-10
@@ -5,7 +5,6 @@ import (
5
"math/rand"
6
"net/http"
7
"slices"
8
- "strings"
8
"time"
9
10
"github.com/gosuda/portal-tunnel/v2/types"
@@ -13,7 +12,6 @@ import (
12
13
type RelayPolicy interface {
14
SelectActive([]RelayState) []RelayState
16
- SelectConfirmed([]RelayState) []RelayState
15
SelectPriority([]RelayState, ClientState) []string
16
OnConfirmed(RelayState) RelayState
17
OnHinted(RelayState) RelayState
@@ -49,12 +47,6 @@ func (p DefaultRelayPolicy) SelectActive(states []RelayState) []RelayState {
47
})
48
}
49
52
-func (p DefaultRelayPolicy) SelectConfirmed(states []RelayState) []RelayState {
53
- return p.selectStates(states, func(state RelayState) bool {
54
- return state.Confirmed
55
- })
56
-}
57
-
50
func (p DefaultRelayPolicy) SelectPriority(states []RelayState, clientState ClientState) []string {
51
selected := p.SelectActive(states)
52
if len(selected) == 0 {
@@ -70,7 +62,7 @@ func (p DefaultRelayPolicy) SelectPriority(states []RelayState, clientState Clie
62
if clientState.RequireTCP && state.hasDescriptor() && !state.Descriptor.SupportsTCP {
63
continue
64
}
73
- relayURL := strings.TrimSpace(state.Descriptor.APIHTTPSAddr)
65
+ relayURL := state.Descriptor.APIHTTPSAddr
66
if slices.Contains(clientState.ExplicitRelayURLs, relayURL) {
67
explicit = append(explicit, relayURL)
68
continue
@@ -86,7 +78,7 @@ func (p DefaultRelayPolicy) SelectPriority(states []RelayState, clientState Clie
78
highRTTAuto := make([]string, 0, len(autoPool))
79
penalizedAuto := make([]string, 0, len(autoPool))
80
for _, state := range autoPool {
89
- relayURL := strings.TrimSpace(state.Descriptor.APIHTTPSAddr)
81
+ relayURL := state.Descriptor.APIHTTPSAddr
82
statePenalized := state.consecutiveFailures > 0 && !state.Reachable
83
switch {
84
case slices.Contains(clientState.ActiveRelayURLs, relayURL) && !statePenalized:
portal/discovery/refresher.go
+37
-43
@@ -62,7 +62,7 @@ func NewRefresher(relaySet *RelaySet, rootCAPEM []byte, overlay OverlayRuntime)
62
}, nil
63
}
64
65
-func (r *Refresher) Refresh(ctx context.Context) error {
65
+func (r *Refresher) Refresh(ctx context.Context, extraSourceURLs ...string) error {
66
if r.overlay != nil {
67
if err := r.refreshOverlay(ctx); err != nil && ctx.Err() == nil {
68
log.Warn().
@@ -73,27 +73,40 @@ func (r *Refresher) Refresh(ctx context.Context) error {
73
return ctx.Err()
74
}
75
}
76
- return r.refreshHTTPS(ctx)
76
+ return r.refreshHTTPS(ctx, extraSourceURLs)
77
}
78
79
-func (r *Refresher) refreshHTTPS(ctx context.Context) error {
79
+func (r *Refresher) refreshHTTPS(ctx context.Context, extraSourceURLs []string) error {
80
r.relaySet.mu.RLock()
81
states := r.relaySet.relayStatesLocked()
82
r.relaySet.mu.RUnlock()
83
84
now := time.Now().UTC()
85
for _, state := range states {
86
- if !state.discoverable(now) || !state.Bootstrap {
86
+ if !state.discoverable(now) || (state.hasDescriptor() && !state.Descriptor.Discovery) {
87
continue
88
}
89
- relay := state.Descriptor
90
- baseURL, err := url.Parse(relay.APIHTTPSAddr)
89
+
90
+ relayURL := state.Descriptor.APIHTTPSAddr
91
+ if relayURL == "" {
92
+ continue
93
+ }
94
+
95
+ recoveryFailures := r.directRecoveryFailures
96
+ if state.Bootstrap {
97
+ recoveryFailures = 0
98
+ }
99
+
100
+ baseURL, err := url.Parse(relayURL)
101
if err != nil {
102
+ if recoveryFailures > 0 {
103
+ r.logDiscoveryFailure(relayURL, relayURL, recoveryFailures, err)
104
+ }
105
continue
106
}
107
if utils.IsLocalRelayHost(baseURL.Hostname()) {
108
log.Info().
96
- Str("relay", relay.APIHTTPSAddr).
109
+ Str("relay", relayURL).
110
Msg("skip loopback relay as discovery source")
111
continue
112
}
@@ -104,57 +117,38 @@ func (r *Refresher) refreshHTTPS(ctx context.Context) error {
117
if ctx.Err() != nil {
118
return ctx.Err()
119
}
120
+ if recoveryFailures > 0 {
121
+ r.logDiscoveryFailure(relayURL, relayURL, recoveryFailures, err)
122
+ }
123
continue
124
}
109
-
125
measuredAt := time.Now().UTC()
111
- if _, err := r.relaySet.ApplyRelayDiscoveryResponse(relay.Identity, relay.APIHTTPSAddr, resp, measuredAt); err != nil {
126
+
127
+ if _, err := r.relaySet.ApplyRelayDiscoveryResponse(relayURL, resp, measuredAt); err != nil {
128
+ if recoveryFailures > 0 {
129
+ r.logDiscoveryFailure(relayURL, relayURL, recoveryFailures, err)
130
+ }
131
continue
132
}
114
- r.relaySet.RecordDiscoveryRTT(relay.APIHTTPSAddr, time.Since(startedAt), measuredAt)
115
- }
116
- if err := ctx.Err(); err != nil {
117
- return err
133
+ r.relaySet.RecordDiscoveryRTT(relayURL, time.Since(startedAt), measuredAt)
134
}
135
120
- r.relaySet.mu.RLock()
121
- states = r.relaySet.relayStatesLocked()
122
- r.relaySet.mu.RUnlock()
123
-
124
- now = time.Now().UTC()
125
- for _, state := range states {
126
- if !state.discoverable(now) || state.Bootstrap {
127
- continue
128
- }
129
- relay := state.Descriptor
130
- baseURL, err := url.Parse(relay.APIHTTPSAddr)
136
+ for _, sourceURL := range extraSourceURLs {
137
+ baseURL, err := url.Parse(sourceURL)
138
if err != nil {
132
- r.logDirectDiscoveryFailure(relay, err, r.directRecoveryFailures)
133
- continue
134
- }
135
- if utils.IsLocalRelayHost(baseURL.Hostname()) {
136
- log.Info().
137
- Str("relay", relay.APIHTTPSAddr).
138
- Msg("skip loopback relay as discovery source")
139
continue
140
}
141
142
- startedAt := time.Now()
142
var resp types.DiscoveryResponse
143
if err := utils.HTTPDoAPIPath(ctx, r.httpClient, baseURL, http.MethodGet, types.PathDiscovery, nil, nil, &resp); err != nil {
144
if ctx.Err() != nil {
145
return ctx.Err()
146
}
148
- r.logDirectDiscoveryFailure(relay, err, r.directRecoveryFailures)
147
continue
148
}
151
-
152
- measuredAt := time.Now().UTC()
153
- if _, err := r.relaySet.ApplyRelayDiscoveryResponse(relay.Identity, relay.APIHTTPSAddr, resp, measuredAt); err != nil {
154
- r.logDirectDiscoveryFailure(relay, err, r.directRecoveryFailures)
149
+ if _, err := r.relaySet.ApplyRelayDiscoveryResponse("", resp, time.Now().UTC()); err != nil {
150
continue
151
}
157
- r.relaySet.RecordDiscoveryRTT(relay.APIHTTPSAddr, time.Since(startedAt), measuredAt)
152
}
153
return nil
154
}
@@ -174,7 +168,7 @@ func (r *Refresher) refreshOverlay(ctx context.Context) error {
168
return err
169
}
170
177
- relaySetChanged, err := r.relaySet.ApplyRelayDiscoveryResponse(relay.Identity, relay.APIHTTPSAddr, resp, time.Now().UTC())
171
+ relaySetChanged, err := r.relaySet.ApplyRelayDiscoveryResponse(relay.APIHTTPSAddr, resp, time.Now().UTC())
172
if err != nil {
173
return err
174
}
@@ -188,19 +182,19 @@ func (r *Refresher) refreshOverlay(ctx context.Context) error {
182
return nil
183
}
184
191
-func (r *Refresher) logDirectDiscoveryFailure(relay types.RelayDescriptor, err error, recoveryFailures int) {
192
- expired, expireReason, consecutiveFailures := r.relaySet.RecordDiscoveryFailure(relay.Identity, relay.APIHTTPSAddr, err, recoveryFailures)
185
+func (r *Refresher) logDiscoveryFailure(targetRelayURL, sourceURL string, recoveryFailures int, err error) {
186
+ expired, expireReason, consecutiveFailures := r.relaySet.RecordRelayFailure(targetRelayURL, err, recoveryFailures)
187
if !expired {
188
return
189
}
190
191
event := log.Warn().
192
Err(err).
199
- Str("relay", relay.APIHTTPSAddr).
193
+ Str("relay", sourceURL).
194
Bool("expired", true).
195
Str("reason", expireReason)
196
if consecutiveFailures > 0 {
197
event = event.Int("consecutive_failures", consecutiveFailures)
198
}
205
- event.Msg("direct relay discovery expired")
199
+ event.Msg("discovery source expired")
200
}
portal/discovery/relayset.go
+92
-142
@@ -3,9 +3,9 @@ package discovery
3
import (
4
"errors"
5
"fmt"
6
+ "net/http"
7
"reflect"
8
"sort"
8
- "strings"
9
"sync"
10
"time"
11
@@ -54,10 +54,7 @@ func (s *RelaySet) SetBootstrapRelayURLs(inputs []string) error {
54
for key, state := range s.relays {
55
_, bootstrap := keep[key]
56
state.Bootstrap = bootstrap
57
- if !state.Bootstrap &&
58
- !state.hasDescriptor() &&
59
- !state.Banned &&
60
- state.consecutiveFailures == 0 {
57
+ if !state.Bootstrap && !state.hasDescriptor() && !state.Banned && state.consecutiveFailures == 0 {
58
delete(s.relays, key)
59
continue
60
}
@@ -115,13 +112,17 @@ func (s *RelaySet) OverlayPeerStates() []RelayState {
112
return out
113
}
114
118
-func (s *RelaySet) ConfirmedDescriptors() []types.RelayDescriptor {
115
+func (s *RelaySet) Descriptors() []types.RelayDescriptor {
116
s.mu.RLock()
120
- defer s.mu.RUnlock()
117
+ states := s.relayStatesLocked()
118
+ s.mu.RUnlock()
119
122
- states := s.policy.SelectConfirmed(s.relayStatesLocked())
120
+ now := time.Now().UTC()
121
out := make([]types.RelayDescriptor, 0, len(states))
122
for _, state := range states {
123
+ if !state.hasDescriptor() || !state.Descriptor.ExpiresAt.After(now) || !state.Descriptor.Discovery {
124
+ continue
125
+ }
126
out = append(out, state.Descriptor)
127
}
128
if len(out) == 0 {
@@ -130,6 +131,40 @@ func (s *RelaySet) ConfirmedDescriptors() []types.RelayDescriptor {
131
return out
132
}
133
134
+func (s *RelaySet) ServeDiscovery(w http.ResponseWriter, r *http.Request, local ...types.RelayDescriptor) {
135
+ if !utils.RequireMethod(w, r, http.MethodGet) {
136
+ return
137
+ }
138
+
139
+ known := s.Descriptors()
140
+ relays := make([]types.RelayDescriptor, 0, len(local)+len(known))
141
+ seen := make(map[string]struct{}, len(local)+len(known))
142
+ add := func(descriptor types.RelayDescriptor) {
143
+ relayURL := descriptor.APIHTTPSAddr
144
+ if relayURL == "" {
145
+ return
146
+ }
147
+ if _, ok := seen[relayURL]; ok {
148
+ return
149
+ }
150
+ seen[relayURL] = struct{}{}
151
+ relays = append(relays, descriptor)
152
+ }
153
+
154
+ for _, descriptor := range local {
155
+ add(descriptor)
156
+ }
157
+ for _, descriptor := range known {
158
+ add(descriptor)
159
+ }
160
+
161
+ utils.WriteAPIData(w, http.StatusOK, types.DiscoveryResponse{
162
+ ProtocolVersion: types.ProtocolVersion,
163
+ GeneratedAt: time.Now().UTC(),
164
+ Relays: relays,
165
+ })
166
+}
167
+
168
func (s *RelaySet) relayStatesLocked() []RelayState {
169
out := make([]RelayState, 0, len(s.relays))
170
for _, state := range s.relays {
@@ -156,123 +191,72 @@ func (s *RelaySet) BanRelayURL(relayURL string) {
191
s.relays[relayURL] = state
192
}
193
159
-func (s *RelaySet) applyDiscoveredStateLocked(state RelayState, confirmed bool) (bool, error) {
160
- relayURL := state.Descriptor.APIHTTPSAddr
161
- relayKey := state.Descriptor.Key()
162
-
163
- previousState, hadPrevious := s.relays[relayURL]
164
- record := previousState
165
- bootstrap := record.Bootstrap
166
-
167
- if hadPrevious && record.hasDescriptor() && record.Descriptor.Key() != relayKey {
168
- return false, errors.New("descriptor identity does not match known relay url")
169
- }
170
-
171
- previousURL := ""
172
- for url, existing := range s.relays {
173
- if url == relayURL || !existing.hasDescriptor() || existing.Descriptor.Key() != relayKey {
174
- continue
175
- }
176
- previousURL = url
177
- bootstrap = bootstrap || existing.Bootstrap
178
- if !record.hasDescriptor() &&
179
- !record.Banned &&
180
- record.consecutiveFailures == 0 {
181
- record.Reachable = existing.Reachable
182
- record.Confirmed = existing.Confirmed
183
- record.Banned = existing.Banned
184
- record.DiscoveryRTT = existing.DiscoveryRTT
185
- record.DiscoveryRTTAt = existing.DiscoveryRTTAt
186
- record.consecutiveFailures = existing.consecutiveFailures
187
- }
188
- break
189
- }
190
-
191
- if !record.hasDescriptor() &&
192
- !record.Banned &&
193
- record.consecutiveFailures == 0 {
194
- record = state
195
- } else {
196
- record.Descriptor = state.Descriptor
197
- record.LastSeenAt = state.LastSeenAt
198
- }
199
- record.Bootstrap = bootstrap
200
-
201
- if confirmed {
202
- record = s.policy.OnConfirmed(record)
203
- } else {
204
- record = s.policy.OnHinted(record)
205
- }
206
-
207
- if previousURL != "" {
208
- delete(s.relays, previousURL)
209
- }
210
- s.relays[relayURL] = record
211
-
212
- return !hadPrevious || previousURL != "" || !reflect.DeepEqual(previousState, record), nil
213
-}
214
-
215
-func (s *RelaySet) ApplyRelayDiscoveryResponse(targetIdentity types.Identity, targetURL string, resp types.DiscoveryResponse, now time.Time) (relaySetChanged bool, err error) {
194
+func (s *RelaySet) ApplyRelayDiscoveryResponse(targetURL string, resp types.DiscoveryResponse, now time.Time) (relaySetChanged bool, err error) {
195
if now.IsZero() {
196
now = time.Now().UTC()
197
} else {
198
now = now.UTC()
199
}
200
222
- s.mu.Lock()
223
- defer s.mu.Unlock()
224
-
201
if resp.ProtocolVersion != types.ProtocolVersion {
202
return false, fmt.Errorf("relay protocol version mismatch: relay=%q client=%q", resp.ProtocolVersion, types.ProtocolVersion)
203
}
204
+ authoritative := targetURL != ""
205
229
- selfState, err := newRelayState(resp.Self, now)
230
- if err != nil {
231
- return false, err
232
- }
233
- if strings.TrimSpace(targetIdentity.Name) == "" && strings.TrimSpace(targetIdentity.Address) == "" {
234
- return false, errors.New("target relay identity is required")
235
- }
236
- if targetName := strings.TrimSpace(targetIdentity.Name); targetName != "" {
237
- if selfState.Descriptor.Name != utils.NormalizeHostname(targetName) {
238
- return false, errors.New("descriptor name does not match target relay")
239
- }
240
- }
241
- if targetAddress := strings.TrimSpace(targetIdentity.Address); targetAddress != "" {
242
- normalizedTargetAddress, err := utils.NormalizeEVMAddress(targetAddress)
243
- if err != nil {
244
- return false, err
245
- }
246
- if selfState.Descriptor.Address != normalizedTargetAddress {
247
- return false, errors.New("descriptor address does not match target relay")
248
- }
249
- }
250
- if targetURL != "" && selfState.Descriptor.APIHTTPSAddr != strings.TrimSpace(targetURL) {
251
- return false, errors.New("descriptor api_https_addr does not match target url")
252
- }
253
- changed, err := s.applyDiscoveredStateLocked(selfState, true)
254
- if err != nil {
255
- return false, err
256
- }
257
- relaySetChanged = relaySetChanged || changed
206
+ s.mu.Lock()
207
+ defer s.mu.Unlock()
208
259
- seen := map[string]struct{}{selfState.Descriptor.Key(): {}}
209
+ discoveredByURL := make(map[string]RelayState, len(resp.Relays))
210
+ discoveredOrder := make([]string, 0, len(resp.Relays))
211
+ targetFound := false
212
for _, descriptor := range resp.Relays {
213
relayState, err := newRelayState(descriptor, now)
214
if err != nil {
215
continue
216
}
265
- relayKey := relayState.Descriptor.Key()
266
- if _, ok := seen[relayKey]; ok {
217
+ relayURL := relayState.Descriptor.APIHTTPSAddr
218
+ if relayURL == "" {
219
continue
220
}
269
- seen[relayKey] = struct{}{}
221
+ if authoritative && relayURL == targetURL {
222
+ targetFound = true
223
+ }
224
+ if _, ok := discoveredByURL[relayURL]; !ok {
225
+ discoveredOrder = append(discoveredOrder, relayURL)
226
+ }
227
+ discoveredByURL[relayURL] = relayState
228
+ }
229
+
230
+ if authoritative && !targetFound {
231
+ return false, errors.New("target relay descriptor missing from relays")
232
+ }
233
+
234
+ for _, relayURL := range discoveredOrder {
235
+ record := discoveredByURL[relayURL]
236
+ existingAtURL, hasExistingAtURL := s.relays[relayURL]
237
+ record.Bootstrap = record.Bootstrap || existingAtURL.Bootstrap
238
+ record.Reachable = record.Reachable || existingAtURL.Reachable
239
+ record.Confirmed = record.Confirmed || existingAtURL.Confirmed
240
+ record.Banned = record.Banned || existingAtURL.Banned
241
+ if record.consecutiveFailures < existingAtURL.consecutiveFailures {
242
+ record.consecutiveFailures = existingAtURL.consecutiveFailures
243
+ }
244
+ if record.DiscoveryRTTAt.IsZero() || (!existingAtURL.DiscoveryRTTAt.IsZero() && existingAtURL.DiscoveryRTTAt.After(record.DiscoveryRTTAt)) {
245
+ record.DiscoveryRTT = existingAtURL.DiscoveryRTT
246
+ record.DiscoveryRTTAt = existingAtURL.DiscoveryRTTAt
247
+ }
248
271
- changed, err := s.applyDiscoveredStateLocked(relayState, false)
272
- if err != nil {
273
- return false, err
249
+ if authoritative && relayURL == targetURL {
250
+ record = s.policy.OnConfirmed(record)
251
+ } else {
252
+ record = s.policy.OnHinted(record)
253
+ }
254
+
255
+ s.relays[relayURL] = record
256
+
257
+ if !hasExistingAtURL || !reflect.DeepEqual(existingAtURL, record) {
258
+ relaySetChanged = true
259
}
275
- relaySetChanged = relaySetChanged || changed
260
}
261
return relaySetChanged, nil
262
}
@@ -291,12 +275,6 @@ func (s *RelaySet) RecordDiscoveryRTT(relayURL string, rtt time.Duration, measur
275
s.relays[relayURL] = state
276
}
277
294
-func (s *RelaySet) recordRelayFailureLocked(relayURL string, state RelayState, err error, recoveryFailures int) (expired bool, expireReason string, consecutiveFailures int) {
295
- state, expired, expireReason = s.policy.OnFailure(state, err, recoveryFailures)
296
- s.relays[relayURL] = state
297
- return expired, expireReason, state.consecutiveFailures
298
-}
299
-
278
func (s *RelaySet) RecordRelayFailure(relayURL string, err error, recoveryFailures int) (expired bool, expireReason string, consecutiveFailures int) {
279
s.mu.Lock()
280
defer s.mu.Unlock()
@@ -305,35 +283,7 @@ func (s *RelaySet) RecordRelayFailure(relayURL string, err error, recoveryFailur
283
if !ok {
284
return false, "", 0
285
}
308
- return s.recordRelayFailureLocked(relayURL, state, err, recoveryFailures)
309
-}
310
-
311
-func (s *RelaySet) RecordDiscoveryFailure(identity types.Identity, relayURL string, err error, recoveryFailures int) (expired bool, expireReason string, consecutiveFailures int) {
312
- relayKey := identity.Key()
313
- if relayKey == "" {
314
- return false, "", 0
315
- }
316
-
317
- s.mu.Lock()
318
- defer s.mu.Unlock()
319
-
320
- state, ok := s.relays[relayURL]
321
- if ok && state.hasDescriptor() && state.Descriptor.Key() != relayKey {
322
- ok = false
323
- }
324
- if !ok {
325
- for url, existing := range s.relays {
326
- if !existing.hasDescriptor() || existing.Descriptor.Key() != relayKey {
327
- continue
328
- }
329
- relayURL = url
330
- state = existing
331
- ok = true
332
- break
333
- }
334
- }
335
- if !ok {
336
- return false, "", 0
337
- }
338
- return s.recordRelayFailureLocked(relayURL, state, err, recoveryFailures)
286
+ state, expired, expireReason = s.policy.OnFailure(state, err, recoveryFailures)
287
+ s.relays[relayURL] = state
288
+ return expired, expireReason, state.consecutiveFailures
289
}
portal/discovery/relayset_test.go
+2
-25
@@ -14,9 +14,9 @@ func TestApplyRelayDiscoveryResponsePreservesBootstrapFlag(t *testing.T) {
14
}
15
16
desc := mustPolicyRelayDescriptor(t, "relay-a", "https://relay-a.example")
17
- if _, err := set.ApplyRelayDiscoveryResponse(desc.Identity, desc.APIHTTPSAddr, types.DiscoveryResponse{
17
+ if _, err := set.ApplyRelayDiscoveryResponse(desc.APIHTTPSAddr, types.DiscoveryResponse{
18
ProtocolVersion: types.ProtocolVersion,
19
- Self: desc,
19
+ Relays: []types.RelayDescriptor{desc},
20
}, time.Now().UTC()); err != nil {
21
t.Fatalf("ApplyRelayDiscoveryResponse() error = %v", err)
22
}
@@ -29,26 +29,3 @@ func TestApplyRelayDiscoveryResponsePreservesBootstrapFlag(t *testing.T) {
29
t.Fatal("bootstrap relay lost bootstrap flag after discovery update")
30
}
31
}
32
-
33
-func TestApplyRelayDiscoveryResponseAllowsURLChangeForSameIdentity(t *testing.T) {
34
- set, err := NewRelaySet(nil)
35
- if err != nil {
36
- t.Fatalf("NewRelaySet() error = %v", err)
37
- }
38
-
39
- desc := mustPolicyRelayDescriptor(t, "relay-a", "https://relay-a.example")
40
- if _, err := set.ApplyRelayDiscoveryResponse(desc.Identity, desc.APIHTTPSAddr, types.DiscoveryResponse{
41
- ProtocolVersion: types.ProtocolVersion,
42
- Self: desc,
43
- }, time.Now().UTC()); err != nil {
44
- t.Fatalf("ApplyRelayDiscoveryResponse() error = %v", err)
45
- }
46
-
47
- changedURL := mustPolicyRelayDescriptor(t, desc.Name, "https://relay-b.example")
48
- if _, err := set.ApplyRelayDiscoveryResponse(desc.Identity, "", types.DiscoveryResponse{
49
- ProtocolVersion: types.ProtocolVersion,
50
- Self: changedURL,
51
- }, time.Now().UTC()); err != nil {
52
- t.Fatalf("ApplyRelayDiscoveryResponse() error = %v, want nil for same relay identity", err)
53
- }
54
-}
portal/discovery/relaystate.go
-13
@@ -2,7 +2,6 @@ package discovery
2
3
import (
4
"errors"
5
- "strings"
5
"time"
6
7
"github.com/gosuda/portal-tunnel/v2/types"
@@ -83,15 +82,3 @@ func (state RelayState) discoverable(now time.Time) bool {
82
}
83
return true
84
}
86
-
87
-func (state RelayState) Equal(other RelayState) bool {
88
- stateKey := state.Descriptor.Key()
89
- otherKey := other.Descriptor.Key()
90
- if stateKey != "" && otherKey != "" && stateKey == otherKey {
91
- return true
92
- }
93
-
94
- stateURL := strings.TrimSpace(state.Descriptor.APIHTTPSAddr)
95
- otherURL := strings.TrimSpace(other.Descriptor.APIHTTPSAddr)
96
- return stateURL != "" && otherURL != "" && stateURL == otherURL
97
-}
portal/lease.go
+77
@@ -4,6 +4,8 @@ import (
4
"context"
5
"errors"
6
"fmt"
7
+ "net"
8
+ "net/url"
9
"strings"
10
"sync"
11
"time"
@@ -292,6 +294,81 @@ func (r *leaseRegistry) countTCPPortLeases() int {
294
return count
295
}
296
297
+func (r *leaseRegistry) snapshot(record *leaseRecord, now time.Time) (types.Lease, bool) {
298
+ if record == nil || now.After(record.ExpiresAt) {
299
+ return types.Lease{}, false
300
+ }
301
+
302
+ adminSnapshot := r.AdminSnapshot(record)
303
+ since := time.Duration(0)
304
+ if !adminSnapshot.LastSeenAt.IsZero() {
305
+ since = max(now.Sub(adminSnapshot.LastSeenAt), 0)
306
+ }
307
+ if adminSnapshot.IsBanned || adminSnapshot.IsDenied || !adminSnapshot.IsApproved || adminSnapshot.Metadata.Hide {
308
+ return types.Lease{}, false
309
+ }
310
+ if adminSnapshot.Ready == 0 && since >= 3*time.Minute {
311
+ return types.Lease{}, false
312
+ }
313
+ return adminSnapshot.Lease, true
314
+}
315
+
316
+func (r *leaseRegistry) LeaseSnapshots(now time.Time) []types.Lease {
317
+ r.mu.RLock()
318
+ defer r.mu.RUnlock()
319
+
320
+ snapshots := make([]types.Lease, 0, len(r.leasesByKey))
321
+ for _, record := range r.leasesByKey {
322
+ snapshot, ok := r.snapshot(record, now)
323
+ if !ok {
324
+ continue
325
+ }
326
+ snapshots = append(snapshots, snapshot)
327
+ }
328
+ return snapshots
329
+}
330
+
331
+func (r *leaseRegistry) AdminLeaseSnapshots(now time.Time) []types.AdminLease {
332
+ r.mu.RLock()
333
+ defer r.mu.RUnlock()
334
+
335
+ snapshots := make([]types.AdminLease, 0, len(r.leasesByKey))
336
+ for _, record := range r.leasesByKey {
337
+ if now.After(record.ExpiresAt) {
338
+ continue
339
+ }
340
+ snapshots = append(snapshots, r.AdminSnapshot(record))
341
+ }
342
+ return snapshots
343
+}
344
+
345
+func (r *leaseRegistry) DiscoverySourceURLs(portalURL string) []string {
346
+ baseURL, err := url.Parse(portalURL)
347
+ if err != nil {
348
+ return nil
349
+ }
350
+
351
+ r.mu.RLock()
352
+ defer r.mu.RUnlock()
353
+
354
+ sources := make([]string, 0, len(r.leasesByKey))
355
+ for _, record := range r.leasesByKey {
356
+ snapshot, ok := r.snapshot(record, time.Now())
357
+ if !ok {
358
+ continue
359
+ }
360
+
361
+ sourceURL := *baseURL
362
+ sourceURL.Host = snapshot.Hostname
363
+ if port := baseURL.Port(); port != "" {
364
+ sourceURL.Host = net.JoinHostPort(snapshot.Hostname, port)
365
+ }
366
+
367
+ sources = append(sources, sourceURL.String())
368
+ }
369
+ return sources
370
+}
371
+
372
func (r *leaseRegistry) Snapshot(record *leaseRecord) types.Lease {
373
snapshot := types.Lease{
374
Name: record.Name,
portal/server.go
+7
-42
@@ -365,52 +365,17 @@ func (s *Server) PortalURL() string {
365
}
366
367
func (s *Server) LeaseSnapshots() []types.Lease {
368
- s.registry.mu.RLock()
369
- defer s.registry.mu.RUnlock()
370
-
371
- now := time.Now()
372
- records := make([]*leaseRecord, 0, len(s.registry.leasesByKey))
373
- for _, record := range s.registry.leasesByKey {
374
- records = append(records, record)
375
- }
376
- snapshots := make([]types.Lease, 0, len(records))
377
- for _, record := range records {
378
- if now.After(record.ExpiresAt) {
379
- continue
380
- }
381
- adminSnapshot := s.registry.AdminSnapshot(record)
382
- since := time.Duration(0)
383
- if !adminSnapshot.LastSeenAt.IsZero() {
384
- since = max(now.Sub(adminSnapshot.LastSeenAt), 0)
385
- }
386
- if adminSnapshot.IsBanned || adminSnapshot.IsDenied || !adminSnapshot.IsApproved || adminSnapshot.Metadata.Hide {
387
- continue
388
- }
389
- if adminSnapshot.Ready == 0 && since >= 3*time.Minute {
390
- continue
391
- }
392
- snapshots = append(snapshots, adminSnapshot.Lease)
368
+ if s == nil || s.registry == nil {
369
+ return nil
370
}
394
- return snapshots
371
+ return s.registry.LeaseSnapshots(time.Now())
372
}
373
374
func (s *Server) AdminLeaseSnapshots() []types.AdminLease {
398
- s.registry.mu.RLock()
399
- defer s.registry.mu.RUnlock()
400
-
401
- now := time.Now()
402
- records := make([]*leaseRecord, 0, len(s.registry.leasesByKey))
403
- for _, record := range s.registry.leasesByKey {
404
- records = append(records, record)
405
- }
406
- snapshots := make([]types.AdminLease, 0, len(records))
407
- for _, record := range records {
408
- if now.After(record.ExpiresAt) {
409
- continue
410
- }
411
- snapshots = append(snapshots, s.registry.AdminSnapshot(record))
375
+ if s == nil || s.registry == nil {
376
+ return nil
377
}
413
- return snapshots
378
+ return s.registry.AdminLeaseSnapshots(time.Now())
379
}
380
381
func (s *Server) LeaseSnapshotByHostname(hostname string) (types.Lease, bool) {
@@ -639,7 +604,7 @@ func (s *Server) runRelayDiscoveryLoop(ctx context.Context) error {
604
defer ticker.Stop()
605
606
for {
642
- if err := refresher.Refresh(ctx); err != nil {
607
+ if err := refresher.Refresh(ctx, s.registry.DiscoverySourceURLs(s.PortalURL())...); err != nil {
608
if ctx.Err() != nil {
609
return nil
610
}
sdk/expose.go
+30
-28
@@ -26,8 +26,8 @@ type Exposure struct {
26
done <-chan struct{}
27
28
identity types.Identity
29
+ discovery bool
30
explicitRelays []string
30
- activeRelayURLs []string
31
TargetAddr string
32
UDPAddr string
33
udpEnabled bool
@@ -115,6 +115,7 @@ func Expose(ctx context.Context, cfg ExposeConfig) (*Exposure, error) {
115
cancel: cancel,
116
done: exposureCtx.Done(),
117
identity: identity,
118
+ discovery: cfg.Discovery,
119
explicitRelays: append([]string(nil), explicitRelayURLs...),
120
TargetAddr: targetAddr,
121
UDPAddr: udpAddr,
@@ -175,20 +176,12 @@ func (e *Exposure) runDiscoveryLoop(ctx context.Context) {
176
func (e *Exposure) ActiveRelayURLs() []string {
177
e.listenerMu.RLock()
178
defer e.listenerMu.RUnlock()
178
- return append([]string(nil), e.activeRelayURLs...)
179
-}
180
-
181
-func (e *Exposure) clientState() discovery.ClientState {
182
- e.listenerMu.RLock()
183
- defer e.listenerMu.RUnlock()
184
-
185
- return discovery.ClientState{
186
- ActiveRelayURLs: append([]string(nil), e.activeRelayURLs...),
187
- ExplicitRelayURLs: append([]string(nil), e.explicitRelays...),
188
- MaxActiveRelays: e.maxActiveRelays,
189
- RequireUDP: e.udpEnabled,
190
- RequireTCP: e.tcpEnabled,
179
+ relayURLs := make([]string, 0, len(e.relayListeners))
180
+ for relayURL := range e.relayListeners {
181
+ relayURLs = append(relayURLs, relayURL)
182
}
183
+ slices.Sort(relayURLs)
184
+ return relayURLs
185
}
186
187
func (e *Exposure) Addr() net.Addr {
@@ -277,6 +270,20 @@ func (e *Exposure) WaitDatagramReady(ctx context.Context) ([]string, error) {
270
}
271
272
func (e *Exposure) RunHTTP(ctx context.Context, handler http.Handler, localAddr string) error {
273
+ if handler == nil {
274
+ handler = http.NotFoundHandler()
275
+ }
276
+ if e.discovery {
277
+ tmp := handler
278
+ handler = http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) {
279
+ if strings.TrimSpace(r.URL.Path) == types.PathDiscovery {
280
+ e.relaySet.ServeDiscovery(w, r)
281
+ return
282
+ }
283
+ tmp.ServeHTTP(w, r)
284
+ })
285
+ }
286
+
287
e.listenerMu.RLock()
288
hasRelayListeners := len(e.relayListeners) > 0
289
e.listenerMu.RUnlock()
@@ -354,7 +361,6 @@ func (e *Exposure) Close() error {
361
e.listenerMu.Lock()
362
relayListeners := e.relayListeners
363
e.relayListeners = make(map[string]*Listener)
357
- e.activeRelayURLs = nil
364
e.listenerMu.Unlock()
365
366
relayURLs := make([]string, 0, len(relayListeners))
@@ -380,7 +386,15 @@ func (e *Exposure) Close() error {
386
}
387
388
func (e *Exposure) reconcileRelayListeners(failOnError bool) error {
383
- desiredRelayURLs := e.relaySet.PriorityRelays(e.clientState())
389
+ clientState := discovery.ClientState{
390
+ ActiveRelayURLs: e.ActiveRelayURLs(),
391
+ ExplicitRelayURLs: append([]string(nil), e.explicitRelays...),
392
+ MaxActiveRelays: e.maxActiveRelays,
393
+ RequireUDP: e.udpEnabled,
394
+ RequireTCP: e.tcpEnabled,
395
+ }
396
+
397
+ desiredRelayURLs := e.relaySet.PriorityRelays(clientState)
398
399
e.listenerMu.Lock()
400
staleRelayListeners := make(map[string]*Listener)
@@ -458,15 +472,6 @@ func (e *Exposure) reconcileRelayListeners(failOnError bool) error {
472
go e.runListenerAcceptLoop(listener)
473
}
474
461
- e.listenerMu.Lock()
462
- e.activeRelayURLs = e.activeRelayURLs[:0]
463
- for _, relayURL := range desiredRelayURLs {
464
- if _, ok := e.relayListeners[relayURL]; !ok {
465
- continue
466
- }
467
- e.activeRelayURLs = append(e.activeRelayURLs, relayURL)
468
- }
469
- e.listenerMu.Unlock()
475
if len(removedRelayURLs) > 0 || len(addedRelayURLs) > 0 {
476
log.Info().
477
Strs("added_relays", addedRelayURLs).
@@ -517,9 +522,6 @@ func (e *Exposure) runListenerAcceptLoop(listener *Listener) {
522
if current, ok := e.relayListeners[relayURL]; ok && current == listener {
523
delete(e.relayListeners, relayURL)
524
}
520
- if index := slices.Index(e.activeRelayURLs, relayURL); index >= 0 {
521
- e.activeRelayURLs = slices.Delete(e.activeRelayURLs, index, index+1)
522
- }
525
e.listenerMu.Unlock()
526
}()
527
sdk/expose_test.go
+2
-3
@@ -33,9 +33,8 @@ func TestExposureReconcileRemovesBannedRelayFromActiveSet(t *testing.T) {
33
}
34
35
exposure := &Exposure{
36
- relaySet: mustRelaySet(t, relayA, relayB),
37
- relayListeners: make(map[string]*Listener, 2),
38
- activeRelayURLs: []string{relayA, relayB},
36
+ relaySet: mustRelaySet(t, relayA, relayB),
37
+ relayListeners: make(map[string]*Listener, 2),
38
}
39
relayAClosed := make(chan struct{})
40
exposure.relayListeners = map[string]*Listener{
types/api.go
+1
-2
@@ -91,8 +91,7 @@ type RegisterResponse struct {
91
type DiscoveryResponse struct {
92
ProtocolVersion string `json:"protocol_version"`
93
GeneratedAt time.Time `json:"generated_at"`
94
- Self RelayDescriptor `json:"self"`
95
- Relays []RelayDescriptor `json:"relays,omitempty"`
94
+ Relays []RelayDescriptor `json:"relays"`
95
}
96
97
type QUICControlMessage struct {
types/identity.go
+1
@@ -114,6 +114,7 @@ type RelayDescriptor struct {
114
WireGuardEndpoint string `json:"wireguard_endpoint,omitempty"`
115
OverlayIPv4 string `json:"overlay_ipv4,omitempty"`
116
OverlayCIDRs []string `json:"overlay_cidrs,omitempty"`
117
+ Discovery bool `json:"discovery,omitempty"`
118
SupportsUDP bool `json:"supports_udp,omitempty"`
119
SupportsTCP bool `json:"supports_tcp,omitempty"`
120
SupportsOverlayPeer bool `json:"supports_overlay_peer,omitempty"`