refact: add refresher and consoliate discovery loop
Kim committed
Apr 9, 2026 at 15:22 UTC
9ffff7bd094065209d9ce1faefcd78db73630073
12 files changed
+815
-1045
portal/discovery/refresher.go
new
+181
@@ -0,0 +1,181 @@
1
+package discovery
2
+
3
+import (
4
+ "context"
5
+ "errors"
6
+ "time"
7
+
8
+ "github.com/rs/zerolog/log"
9
+
10
+ "github.com/gosuda/portal-tunnel/v2/types"
11
+)
12
+
13
+const (
14
+ defaultRecoveryFailures = 3
15
+)
16
+
17
+type OverlayRuntime interface {
18
+ DiscoverRelay(context.Context, types.RelayDescriptor) (types.DiscoveryResponse, error)
19
+ Sync(map[string]RelayState) error
20
+}
21
+
22
+type Refresher struct {
23
+ relaySet *RelaySet
24
+ rootCAPEM []byte
25
+ overlay OverlayRuntime
26
+ directRecoveryFailures int
27
+ overlayRecoveryFailures int
28
+}
29
+
30
+func NewRefresher(relaySet *RelaySet, rootCAPEM []byte, overlay OverlayRuntime) (*Refresher, error) {
31
+ if relaySet == nil {
32
+ return nil, errors.New("relay set is required")
33
+ }
34
+ return &Refresher{
35
+ relaySet: relaySet,
36
+ rootCAPEM: append([]byte(nil), rootCAPEM...),
37
+ overlay: overlay,
38
+ directRecoveryFailures: defaultRecoveryFailures,
39
+ overlayRecoveryFailures: defaultRecoveryFailures,
40
+ }, nil
41
+}
42
+
43
+func (r *Refresher) Refresh(ctx context.Context) error {
44
+ if err := r.refreshHTTPS(ctx); err != nil {
45
+ return err
46
+ }
47
+ if r.overlay == nil {
48
+ return ctx.Err()
49
+ }
50
+ if err := r.overlay.Sync(r.relaySet.View()); err != nil {
51
+ log.Warn().
52
+ Err(err).
53
+ Msg("sync wireguard peers")
54
+ return ctx.Err()
55
+ }
56
+ return r.refreshOverlay(ctx)
57
+}
58
+
59
+func (r *Refresher) refreshHTTPS(ctx context.Context) error {
60
+ for _, bootstrap := range r.relaySet.BootstrapDescriptors() {
61
+ resp, err := DiscoverRelayDiscovery(ctx, bootstrap.APIHTTPSAddr, r.rootCAPEM, nil)
62
+ if err != nil {
63
+ if ctx.Err() != nil {
64
+ return ctx.Err()
65
+ }
66
+ if _, _, unavailable := DiscoveryUnavailableStatus(err); unavailable {
67
+ continue
68
+ }
69
+ continue
70
+ }
71
+
72
+ now := time.Now().UTC()
73
+ _, _, err = r.relaySet.ApplyRelayDiscoveryResponse(bootstrap.Identity, bootstrap.APIHTTPSAddr, resp, now)
74
+ if err != nil {
75
+ continue
76
+ }
77
+ }
78
+ if err := ctx.Err(); err != nil {
79
+ return err
80
+ }
81
+
82
+ for _, relay := range r.relaySet.confirmableDescriptors() {
83
+ if r.overlay != nil && relay.SupportsOverlayPeer {
84
+ continue
85
+ }
86
+ resp, err := DiscoverRelayDiscovery(ctx, relay.APIHTTPSAddr, r.rootCAPEM, nil)
87
+ if err != nil {
88
+ if ctx.Err() != nil {
89
+ return ctx.Err()
90
+ }
91
+ r.logDirectDiscoveryFailure(relay, err, r.directRecoveryFailures)
92
+ continue
93
+ }
94
+
95
+ now := time.Now().UTC()
96
+ _, _, err = r.relaySet.ApplyRelayDiscoveryResponse(relay.Identity, relay.APIHTTPSAddr, resp, now)
97
+ if err != nil {
98
+ r.logDirectDiscoveryFailure(relay, err, r.directRecoveryFailures)
99
+ continue
100
+ }
101
+ }
102
+ return ctx.Err()
103
+}
104
+
105
+func (r *Refresher) refreshOverlay(ctx context.Context) error {
106
+ for _, relay := range r.relaySet.SyncableDescriptors() {
107
+ var failureErr error
108
+
109
+ if err := RequireOverlayRelayDescriptor(relay); err != nil {
110
+ failureErr = err
111
+ } else {
112
+ resp, err := r.overlay.DiscoverRelay(ctx, relay)
113
+ if err != nil {
114
+ if ctx.Err() != nil {
115
+ return ctx.Err()
116
+ }
117
+ failureErr = err
118
+ } else {
119
+ now := time.Now().UTC()
120
+ relaySetChanged, warnErr, err := r.relaySet.ApplyOverlayRelayDiscoveryResponse(relay.Identity, relay.APIHTTPSAddr, resp, now)
121
+ if relaySetChanged {
122
+ view := r.relaySet.View()
123
+ if syncErr := r.overlay.Sync(view); syncErr != nil {
124
+ if warnErr == nil {
125
+ warnErr = syncErr
126
+ }
127
+ }
128
+ }
129
+ if err != nil {
130
+ failureErr = err
131
+ } else {
132
+ if warnErr != nil {
133
+ log.Warn().
134
+ Err(warnErr).
135
+ Str("relay", relay.APIHTTPSAddr).
136
+ Msg("overlay relay discovery completed with warnings")
137
+ }
138
+ continue
139
+ }
140
+ }
141
+ }
142
+
143
+ expired, expireReason, consecutiveFailures := r.relaySet.RecordDiscoveryFailure(relay.Identity, relay.APIHTTPSAddr, failureErr, r.overlayRecoveryFailures)
144
+ if expired {
145
+ if syncErr := r.overlay.Sync(r.relaySet.View()); syncErr != nil && failureErr == nil {
146
+ failureErr = syncErr
147
+ }
148
+ }
149
+
150
+ event := log.Warn().
151
+ Err(failureErr).
152
+ Str("relay", relay.APIHTTPSAddr)
153
+ if expired {
154
+ event = event.
155
+ Bool("expired", true).
156
+ Str("reason", expireReason)
157
+ if consecutiveFailures > 0 {
158
+ event = event.Int("consecutive_failures", consecutiveFailures)
159
+ }
160
+ }
161
+ event.Msg("overlay relay discovery failed")
162
+ }
163
+ return ctx.Err()
164
+}
165
+
166
+func (r *Refresher) logDirectDiscoveryFailure(relay types.RelayDescriptor, err error, recoveryFailures int) {
167
+ expired, expireReason, consecutiveFailures := r.relaySet.RecordDiscoveryFailure(relay.Identity, relay.APIHTTPSAddr, err, recoveryFailures)
168
+ if !expired {
169
+ return
170
+ }
171
+
172
+ event := log.Warn().
173
+ Err(err).
174
+ Str("relay", relay.APIHTTPSAddr).
175
+ Bool("expired", true).
176
+ Str("reason", expireReason)
177
+ if consecutiveFailures > 0 {
178
+ event = event.Int("consecutive_failures", consecutiveFailures)
179
+ }
180
+ event.Msg("direct relay discovery expired")
181
+}
portal/discovery/relayset.go
+284
-726
@@ -1,130 +1,108 @@
1
package discovery
2
3
import (
4
- "context"
4
"errors"
5
"net/http"
6
"reflect"
7
+ "slices"
8
"sort"
9
"strings"
10
"sync"
11
"time"
12
13
- "github.com/rs/zerolog/log"
14
-
13
"github.com/gosuda/portal-tunnel/v2/types"
14
"github.com/gosuda/portal-tunnel/v2/utils"
15
)
16
19
-type RelayView struct {
20
- Descriptor types.RelayDescriptor
21
- FirstSeenAt time.Time
22
- LastSeenAt time.Time
23
-}
17
+type relayStatus uint8
18
+
19
+const (
20
+ relayStatusHinted relayStatus = iota
21
+ relayStatusConfirmed
22
+ relayStatusExpired
23
+)
24
25
-type RelayLocalState struct {
25
+// RelayState is the single relay shape shared by discovery storage and overlay sync.
26
+// Package-private fields are internal-only local state.
27
+type RelayState struct {
28
+ Descriptor types.RelayDescriptor
29
+ FirstSeenAt time.Time
30
+ LastSeenAt time.Time
31
Banned bool
27
- BanReason string
28
- Bootstrap bool
29
- Advertised bool
32
Expired bool
31
- Reachable bool
32
- ConsecutiveFailures int
33
- LastSuccessAt time.Time
34
- LastFailureAt time.Time
33
+ status relayStatus
34
+ consecutiveFailures int
35
}
36
37
-type RelaySummary struct {
38
- Known int
39
- Banned int
40
- Bootstrap int
41
- Advertised int
42
- Expired int
43
- Syncable int
44
- Reachable int
45
- Unreachable int
37
+func (s RelayState) isDefaultLocalState() bool {
38
+ return !s.Banned && s.status == relayStatusHinted && s.consecutiveFailures == 0
39
}
40
48
-// RelaySet owns the shared relay discovery view: known relay URLs, stable relay
49
-// id/url mappings, the latest validated descriptor seen for each relay, and common
50
-// process-local relay state such as ban/reachability/failure tracking.
51
-//
52
-// Runtime-specific policy such as bootstrap classification, relay lifecycle, or
53
-// listener ownership belongs in the caller's projection.
41
+// RelaySet owns the shared relay discovery view: configured bootstrap relay URLs,
42
+// the latest validated descriptor seen for each relay, and local runtime state
43
+// such as ban/reachability/failure tracking.
44
type RelaySet struct {
55
- mu sync.RWMutex
56
- knownRelayURLs []string
57
- relayKeysByURL map[string]string
58
- relays map[string]RelayView
59
- localByURL map[string]RelayLocalState
60
- lastStatusReachable map[string]bool
61
- lastStatusSummary RelaySummary
62
- haveLastStatus bool
63
- selfRelayKey string
64
- selfRelayURL string
45
+ mu sync.RWMutex
46
+ knownRelayURLs []string
47
+ relayKeysByURL map[string]string
48
+ relays map[string]RelayState
49
+ localByURL map[string]RelayState
50
+ selfRelayKey string
51
+ selfRelayURL string
52
}
53
67
-const defaultDiscoveryRecoveryFailures = 3
54
+type relayDescriptorProjection struct {
55
+ state RelayState
56
+ relayURL string
57
+ bootstrap bool
58
+}
59
60
func NewRelaySet() *RelaySet {
61
return &RelaySet{
62
relayKeysByURL: make(map[string]string),
72
- relays: make(map[string]RelayView),
73
- localByURL: make(map[string]RelayLocalState),
63
+ relays: make(map[string]RelayState),
64
+ localByURL: make(map[string]RelayState),
65
}
66
}
67
77
-func (s *RelaySet) isSelfRelayURLLocked(relayURL string) bool {
78
- if s == nil {
68
+func relayExpiredAt(state RelayState, now time.Time) bool {
69
+ if state.status == relayStatusExpired {
70
+ return true
71
+ }
72
+ if state.Descriptor.ExpiresAt.IsZero() {
73
return false
74
}
75
+ if now.IsZero() {
76
+ now = time.Now().UTC()
77
+ }
78
+ return !state.Descriptor.ExpiresAt.After(now)
79
+}
80
+
81
+func (s *RelaySet) isSelfRelayURLLocked(relayURL string) bool {
82
relayURL = strings.TrimSpace(relayURL)
83
return relayURL != "" && s.selfRelayURL != "" && relayURL == s.selfRelayURL
84
}
85
86
func (s *RelaySet) isSelfRelayDescriptorLocked(desc types.RelayDescriptor) bool {
86
- if s == nil {
87
- return false
88
- }
89
- if s.selfRelayKey != "" {
90
- if relayKey := desc.Key(); relayKey != "" && relayKey == s.selfRelayKey {
91
- return true
92
- }
87
+ if relayKey := desc.Key(); relayKey != "" && s.selfRelayKey != "" && relayKey == s.selfRelayKey {
88
+ return true
89
}
90
return s.isSelfRelayURLLocked(desc.APIHTTPSAddr)
91
}
92
97
-func (s *RelaySet) pruneSelfRelayLocked() {
98
- if s == nil {
93
+func (s *RelaySet) storeLocalStateLocked(relayURL string, state RelayState) {
94
+ relayURL = strings.TrimSpace(relayURL)
95
+ if relayURL == "" {
96
return
97
}
101
- if s.selfRelayURL != "" {
102
- filtered := s.knownRelayURLs[:0]
103
- for _, relayURL := range s.knownRelayURLs {
104
- if s.isSelfRelayURLLocked(relayURL) {
105
- continue
106
- }
107
- filtered = append(filtered, relayURL)
108
- }
109
- s.knownRelayURLs = filtered
110
- delete(s.localByURL, s.selfRelayURL)
111
- delete(s.relayKeysByURL, s.selfRelayURL)
112
- }
113
- if s.selfRelayKey != "" {
114
- if view, ok := s.relays[s.selfRelayKey]; ok {
115
- delete(s.localByURL, view.Descriptor.APIHTTPSAddr)
116
- delete(s.relayKeysByURL, view.Descriptor.APIHTTPSAddr)
117
- }
118
- delete(s.relays, s.selfRelayKey)
98
+ if state.isDefaultLocalState() {
99
+ delete(s.localByURL, relayURL)
100
+ return
101
}
102
+ s.localByURL[relayURL] = state
103
}
104
122
-// SetSelfRelay configures the relay set with the local relay identity so it
123
-// can skip self-references when evaluating discovery hints.
105
func (s *RelaySet) SetSelfRelay(identity types.Identity, relayURL string) error {
125
- if s == nil {
126
- return nil
127
- }
106
relayURL = strings.TrimSpace(relayURL)
107
if relayURL != "" {
108
normalized, err := utils.NormalizeRelayURL(relayURL)
@@ -138,56 +116,36 @@ func (s *RelaySet) SetSelfRelay(identity types.Identity, relayURL string) error
116
defer s.mu.Unlock()
117
s.selfRelayKey = identity.Key()
118
s.selfRelayURL = relayURL
141
- s.pruneSelfRelayLocked()
142
- s.logStatusChange()
143
- return nil
144
-}
145
-
146
-func (s *RelaySet) trackedRelayURLs() []string {
147
- if s == nil {
148
- return nil
149
- }
150
-
151
- urls := make([]string, 0, len(s.knownRelayURLs)+len(s.relays))
152
- seen := make(map[string]struct{}, len(s.knownRelayURLs)+len(s.relays))
153
- for _, relayURL := range s.knownRelayURLs {
154
- relayURL = strings.TrimSpace(relayURL)
155
- if relayURL == "" {
156
- continue
157
- }
158
- if _, ok := seen[relayURL]; ok {
159
- continue
119
+ if s.selfRelayURL != "" {
120
+ filtered := s.knownRelayURLs[:0]
121
+ for _, knownRelayURL := range s.knownRelayURLs {
122
+ if s.isSelfRelayURLLocked(knownRelayURL) {
123
+ continue
124
+ }
125
+ filtered = append(filtered, knownRelayURL)
126
}
161
- seen[relayURL] = struct{}{}
162
- urls = append(urls, relayURL)
127
+ s.knownRelayURLs = filtered
128
+ delete(s.localByURL, s.selfRelayURL)
129
+ delete(s.relayKeysByURL, s.selfRelayURL)
130
}
164
- for _, view := range s.relays {
165
- relayURL := view.Descriptor.APIHTTPSAddr
166
- if relayURL == "" {
167
- continue
168
- }
169
- if _, ok := seen[relayURL]; ok {
170
- continue
131
+ if s.selfRelayKey != "" {
132
+ if record, ok := s.relays[s.selfRelayKey]; ok {
133
+ delete(s.localByURL, record.Descriptor.APIHTTPSAddr)
134
+ delete(s.relayKeysByURL, record.Descriptor.APIHTTPSAddr)
135
}
172
- seen[relayURL] = struct{}{}
173
- urls = append(urls, relayURL)
136
+ delete(s.relays, s.selfRelayKey)
137
}
175
- return urls
138
+ return nil
139
}
140
178
-func (s *RelaySet) ActiveRelayURLs() []string {
179
- if s == nil {
180
- return nil
181
- }
182
- s.mu.RLock()
183
- defer s.mu.RUnlock()
141
+func (s *RelaySet) bootstrapRelayURLsLocked() []string {
142
if len(s.knownRelayURLs) == 0 {
143
return nil
144
}
145
146
out := make([]string, 0, len(s.knownRelayURLs))
147
for _, relayURL := range s.knownRelayURLs {
190
- if state, ok := s.localByURL[relayURL]; ok && state.Banned {
148
+ if s.isSelfRelayURLLocked(relayURL) || s.localByURL[relayURL].Banned {
149
continue
150
}
151
out = append(out, relayURL)
@@ -198,121 +156,95 @@ func (s *RelaySet) ActiveRelayURLs() []string {
156
return out
157
}
158
201
-func relayExpiredAt(view RelayView, state RelayLocalState, now time.Time) bool {
202
- if state.Expired {
203
- return true
204
- }
205
- if view.Descriptor.ExpiresAt.IsZero() {
206
- return false
207
- }
208
- if now.IsZero() {
209
- now = time.Now().UTC()
159
+func (s *RelaySet) descriptorProjectionsLocked() []relayDescriptorProjection {
160
+ if len(s.relays) == 0 {
161
+ return nil
162
}
211
- return !view.Descriptor.ExpiresAt.After(now)
212
-}
163
214
-func (s *RelaySet) logStatusChange() {
215
- now := time.Now().UTC()
216
- var currentReachable map[string]bool
217
- trackedRelayURLs := s.trackedRelayURLs()
218
- if len(trackedRelayURLs) > 0 {
219
- currentReachable = make(map[string]bool, len(trackedRelayURLs))
220
- for _, relayURL := range trackedRelayURLs {
221
- state := s.localByURL[relayURL]
222
- currentReachable[relayURL] = !state.Banned && state.Reachable
223
- }
224
- }
225
- summary := RelaySummary{}
226
- for _, relayURL := range trackedRelayURLs {
227
- summary.Known++
228
- state := s.localByURL[relayURL]
229
- relayKey := s.relayKeysByURL[relayURL]
230
- view, ok := s.relays[relayKey]
231
- expired := ok && relayExpiredAt(view, state, now) || !ok && state.Expired
232
- if state.Banned {
233
- summary.Banned++
164
+ out := make([]relayDescriptorProjection, 0, len(s.relays))
165
+ for _, record := range s.relays {
166
+ if s.isSelfRelayDescriptorLocked(record.Descriptor) {
167
continue
168
}
236
- if state.Bootstrap {
237
- summary.Bootstrap++
238
- }
239
- if state.Advertised && !expired {
240
- summary.Advertised++
241
- }
242
- if expired {
243
- summary.Expired++
244
- }
245
- if state.Reachable {
246
- summary.Reachable++
247
- } else {
248
- summary.Unreachable++
169
+ relayURL := strings.TrimSpace(record.Descriptor.APIHTTPSAddr)
170
+ if relayURL == "" {
171
+ continue
172
}
250
- if ok && !state.Bootstrap && !expired && view.Descriptor.SupportsOverlayPeer {
251
- summary.Syncable++
173
+ bootstrap := false
174
+ if !s.isSelfRelayURLLocked(relayURL) {
175
+ for _, candidate := range s.knownRelayURLs {
176
+ if candidate == relayURL {
177
+ bootstrap = true
178
+ break
179
+ }
180
+ }
181
}
182
+ local := s.localByURL[relayURL]
183
+ record.Banned = local.Banned
184
+ record.status = local.status
185
+ record.consecutiveFailures = local.consecutiveFailures
186
+ out = append(out, relayDescriptorProjection{
187
+ state: record,
188
+ relayURL: relayURL,
189
+ bootstrap: bootstrap,
190
+ })
191
}
254
- if s.haveLastStatus && summary == s.lastStatusSummary && reflect.DeepEqual(currentReachable, s.lastStatusReachable) {
255
- return
192
+ if len(out) == 0 {
193
+ return nil
194
}
195
+ sort.Slice(out, func(i, j int) bool {
196
+ return out[i].relayURL < out[j].relayURL
197
+ })
198
+ return out
199
+}
200
+
201
+func (s *RelaySet) ActiveRelayURLs() []string {
202
+ s.mu.RLock()
203
+ defer s.mu.RUnlock()
204
+
205
+ now := time.Now().UTC()
206
+ bootstrapRelayURLs := s.bootstrapRelayURLsLocked()
207
+ projections := s.descriptorProjectionsLocked()
208
258
- activated := make([]string, 0)
259
- deactivated := make([]string, 0)
260
- for relayURL, reachable := range currentReachable {
261
- if s.lastStatusReachable == nil || s.lastStatusReachable[relayURL] == reachable {
209
+ out := make([]string, 0, len(bootstrapRelayURLs)+len(projections))
210
+ seen := make(map[string]struct{}, len(bootstrapRelayURLs)+len(projections))
211
+ for _, relayURL := range bootstrapRelayURLs {
212
+ if _, ok := seen[relayURL]; ok {
213
continue
214
}
264
- if reachable {
265
- activated = append(activated, relayURL)
266
- } else {
267
- deactivated = append(deactivated, relayURL)
268
- }
215
+ seen[relayURL] = struct{}{}
216
+ out = append(out, relayURL)
217
}
270
- for relayURL, reachable := range s.lastStatusReachable {
271
- if _, ok := currentReachable[relayURL]; ok || !reachable {
218
+ for _, projection := range projections {
219
+ if projection.state.Banned || projection.state.status != relayStatusConfirmed || relayExpiredAt(projection.state, now) {
220
continue
221
}
274
- deactivated = append(deactivated, relayURL)
275
- }
276
-
277
- event := log.Info().
278
- Int("banned", summary.Banned).
279
- Int("bootstrap", summary.Bootstrap).
280
- Int("advertised", summary.Advertised).
281
- Int("expired", summary.Expired).
282
- Int("syncable", summary.Syncable).
283
- Int("reachable", summary.Reachable).
284
- Int("unreachable", summary.Unreachable)
285
- if len(activated) > 0 {
286
- event = event.Strs("activated", activated)
222
+ if _, ok := seen[projection.relayURL]; ok {
223
+ continue
224
+ }
225
+ seen[projection.relayURL] = struct{}{}
226
+ out = append(out, projection.relayURL)
227
}
288
- if len(deactivated) > 0 {
289
- event = event.Strs("deactivated", deactivated)
228
+ if len(out) == 0 {
229
+ return nil
230
}
291
- event.Msg("relay status")
292
- s.lastStatusReachable = currentReachable
293
- s.lastStatusSummary = summary
294
- s.haveLastStatus = true
231
+ return out
232
}
233
234
func (s *RelaySet) BootstrapDescriptors() []types.RelayDescriptor {
298
- if s == nil {
299
- return nil
300
- }
235
s.mu.RLock()
236
defer s.mu.RUnlock()
303
- if len(s.knownRelayURLs) == 0 {
237
+
238
+ bootstrapRelayURLs := s.bootstrapRelayURLsLocked()
239
+ if len(bootstrapRelayURLs) == 0 {
240
return nil
241
}
242
307
- out := make([]types.RelayDescriptor, 0, len(s.knownRelayURLs))
308
- for _, relayURL := range s.knownRelayURLs {
309
- state, ok := s.localByURL[relayURL]
310
- if !ok || !state.Bootstrap {
311
- continue
312
- }
243
+ out := make([]types.RelayDescriptor, 0, len(bootstrapRelayURLs))
244
+ for _, relayURL := range bootstrapRelayURLs {
245
if relayKey, ok := s.relayKeysByURL[relayURL]; ok {
314
- if view, ok := s.relays[relayKey]; ok && view.Descriptor.APIHTTPSAddr != "" {
315
- out = append(out, view.Descriptor)
246
+ if record, ok := s.relays[relayKey]; ok && record.Descriptor.APIHTTPSAddr != "" {
247
+ out = append(out, record.Descriptor)
248
continue
249
}
250
}
@@ -320,6 +252,7 @@ func (s *RelaySet) BootstrapDescriptors() []types.RelayDescriptor {
252
Identity: types.Identity{
253
Name: utils.PortalRootHost(relayURL),
254
},
255
+ RelayID: relayURL,
256
APIHTTPSAddr: relayURL,
257
Version: 1,
258
})
@@ -330,362 +263,152 @@ func (s *RelaySet) BootstrapDescriptors() []types.RelayDescriptor {
263
return out
264
}
265
333
-func (s *RelaySet) ActiveRelayDescriptors() []types.RelayDescriptor {
334
- advertised := s.AdvertisedDescriptors()
335
- if len(advertised) == 0 {
336
- return nil
337
- }
338
- s.mu.RLock()
339
- defer s.mu.RUnlock()
340
-
341
- filtered := advertised[:0]
342
- for _, desc := range advertised {
343
- if s.isSelfRelayDescriptorLocked(desc) {
344
- continue
345
- }
346
- filtered = append(filtered, desc)
347
- }
348
- if len(filtered) == 0 {
349
- return nil
350
- }
351
- return append([]types.RelayDescriptor(nil), filtered...)
352
-}
353
-
354
-func (s *RelaySet) BanRelayURL(relayURL, reason string) bool {
355
- if s == nil {
356
- return false
357
- }
266
+func (s *RelaySet) BanRelayURL(relayURL string) {
267
s.mu.Lock()
268
defer s.mu.Unlock()
269
relayURL = strings.TrimSpace(relayURL)
270
if relayURL == "" {
362
- return false
271
+ return
272
}
273
274
state := s.localByURL[relayURL]
366
- reason = strings.TrimSpace(reason)
367
- changed := !state.Banned || strings.TrimSpace(state.BanReason) != reason
275
state.Banned = true
369
- state.BanReason = reason
370
- state.Reachable = false
371
- s.localByURL[relayURL] = state
372
- if changed {
373
- s.logStatusChange()
374
- }
375
- return changed
276
+ s.storeLocalStateLocked(relayURL, state)
277
}
278
378
-func (s *RelaySet) MarkRelayUnreachable(relayURL string) bool {
379
- if s == nil {
380
- return false
381
- }
382
- s.mu.Lock()
383
- defer s.mu.Unlock()
384
- relayURL = strings.TrimSpace(relayURL)
385
- if relayURL == "" {
386
- return false
387
- }
388
-
389
- state := s.localByURL[relayURL]
390
- if state.Banned {
391
- return false
392
- }
393
- if !state.Reachable {
394
- return false
395
- }
396
- state.Reachable = false
397
- s.localByURL[relayURL] = state
398
- s.logStatusChange()
399
- return true
400
-}
401
-
402
-func (s *RelaySet) MarkRelayReachable(relayURL string, now time.Time) bool {
403
- if s == nil {
404
- return false
405
- }
406
- s.mu.Lock()
407
- defer s.mu.Unlock()
408
- relayURL = strings.TrimSpace(relayURL)
409
- if relayURL == "" {
410
- return false
411
- }
412
- if now.IsZero() {
413
- now = time.Now().UTC()
414
- }
415
-
416
- state := s.localByURL[relayURL]
417
- changed := !state.Reachable || state.ConsecutiveFailures != 0 || state.LastSuccessAt != now
418
- state.Reachable = true
419
- state.ConsecutiveFailures = 0
420
- state.LastSuccessAt = now
421
- s.localByURL[relayURL] = state
422
- if changed {
423
- s.logStatusChange()
424
- }
425
- return changed
426
-}
279
+func (s *RelaySet) AdvertisedDescriptors() []types.RelayDescriptor {
280
+ s.mu.RLock()
281
+ defer s.mu.RUnlock()
282
428
-func (s *RelaySet) MarkRelayFailure(relayURL string, now time.Time) RelayLocalState {
429
- if s == nil {
430
- return RelayLocalState{}
431
- }
432
- s.mu.Lock()
433
- defer s.mu.Unlock()
434
- relayURL = strings.TrimSpace(relayURL)
435
- if relayURL == "" {
436
- return RelayLocalState{}
437
- }
438
- if now.IsZero() {
439
- now = time.Now().UTC()
283
+ now := time.Now().UTC()
284
+ projections := s.descriptorProjectionsLocked()
285
+ if len(projections) == 0 {
286
+ return nil
287
}
288
442
- state := s.localByURL[relayURL]
443
- state.Reachable = false
444
- state.ConsecutiveFailures++
445
- state.LastFailureAt = now
446
- s.localByURL[relayURL] = state
447
- s.logStatusChange()
448
- return state
449
-}
450
-
451
-func (s *RelaySet) RecordBootstrapDiscoveryFailure(relayURL string, err error, now time.Time) {
452
- state := s.MarkRelayFailure(relayURL, now)
453
- if statusCode, code, unavailable := DiscoveryUnavailableStatus(err); unavailable {
454
- if state.ConsecutiveFailures > 1 {
455
- return
456
- }
457
- event := log.Info().Str("relay", relayURL)
458
- if statusCode > 0 {
459
- event = event.Int("status_code", statusCode)
460
- }
461
- if code != "" {
462
- event = event.Str("code", code)
289
+ out := make([]types.RelayDescriptor, 0, len(projections))
290
+ for _, projection := range projections {
291
+ if projection.state.Banned || projection.state.status != relayStatusConfirmed || relayExpiredAt(projection.state, now) {
292
+ continue
293
}
464
- event.Msg("bootstrap relay discovery unavailable; peer may have discovery disabled")
465
- return
294
+ out = append(out, projection.state.Descriptor)
295
}
467
-
468
- log.Warn().
469
- Err(err).
470
- Str("relay", relayURL).
471
- Msg("bootstrap relay discovery failed")
472
-}
473
-
474
-func (s *RelaySet) AdvertisedDescriptors() []types.RelayDescriptor {
475
- if s == nil {
296
+ if len(out) == 0 {
297
return nil
298
}
299
+ return out
300
+}
301
+
302
+func (s *RelaySet) confirmableDescriptors() []types.RelayDescriptor {
303
s.mu.RLock()
304
defer s.mu.RUnlock()
480
- if len(s.relays) == 0 {
305
+
306
+ now := time.Now().UTC()
307
+ projections := s.descriptorProjectionsLocked()
308
+ if len(projections) == 0 {
309
return nil
310
}
311
484
- now := time.Now().UTC()
485
- out := make([]types.RelayDescriptor, 0, len(s.relays))
486
- for _, view := range s.relays {
487
- state := s.localByURL[view.Descriptor.APIHTTPSAddr]
488
- if !state.Advertised || relayExpiredAt(view, state, now) || view.Descriptor.APIHTTPSAddr == "" {
312
+ out := make([]types.RelayDescriptor, 0, len(projections))
313
+ for _, projection := range projections {
314
+ if projection.bootstrap || projection.state.Banned || relayExpiredAt(projection.state, now) {
315
continue
316
}
491
- out = append(out, view.Descriptor)
317
+ out = append(out, projection.state.Descriptor)
318
}
319
if len(out) == 0 {
320
return nil
321
}
496
- sort.Slice(out, func(i, j int) bool {
497
- return out[i].APIHTTPSAddr < out[j].APIHTTPSAddr
498
- })
322
return out
323
}
324
325
func (s *RelaySet) SyncableDescriptors() []types.RelayDescriptor {
503
- if s == nil {
504
- return nil
505
- }
326
s.mu.RLock()
327
defer s.mu.RUnlock()
508
- if len(s.relays) == 0 {
328
+
329
+ now := time.Now().UTC()
330
+ projections := s.descriptorProjectionsLocked()
331
+ if len(projections) == 0 {
332
return nil
333
}
334
512
- now := time.Now().UTC()
513
- out := make([]types.RelayDescriptor, 0, len(s.relays))
514
- for _, view := range s.relays {
515
- state := s.localByURL[view.Descriptor.APIHTTPSAddr]
516
- if state.Bootstrap || relayExpiredAt(view, state, now) || !view.Descriptor.SupportsOverlayPeer {
335
+ out := make([]types.RelayDescriptor, 0, len(projections))
336
+ for _, projection := range projections {
337
+ if projection.bootstrap || projection.state.Banned || relayExpiredAt(projection.state, now) || !projection.state.Descriptor.SupportsOverlayPeer {
338
continue
339
}
519
- out = append(out, view.Descriptor)
340
+ out = append(out, projection.state.Descriptor)
341
}
342
if len(out) == 0 {
343
return nil
344
}
524
- sort.Slice(out, func(i, j int) bool {
525
- return out[i].APIHTTPSAddr < out[j].APIHTTPSAddr
526
- })
345
return out
346
}
347
530
-func (s *RelaySet) Snapshot() map[string]types.RelayState {
531
- if s == nil {
532
- return nil
533
- }
348
+func (s *RelaySet) View() map[string]RelayState {
349
s.mu.RLock()
350
defer s.mu.RUnlock()
536
- if len(s.relays) == 0 {
537
- return nil
538
- }
351
352
now := time.Now().UTC()
541
- snapshot := make(map[string]types.RelayState, len(s.relays))
542
- for relayKey, view := range s.relays {
543
- localState := s.localByURL[view.Descriptor.APIHTTPSAddr]
544
- snapshot[relayKey] = types.RelayState{
545
- Descriptor: view.Descriptor,
546
- Bootstrap: localState.Bootstrap,
547
- Advertised: localState.Advertised,
548
- Expired: relayExpiredAt(view, localState, now),
549
- FirstSeenAt: view.FirstSeenAt,
550
- LastSeenAt: view.LastSeenAt,
551
- ConsecutiveFailures: localState.ConsecutiveFailures,
552
- }
553
- }
554
- return snapshot
555
-}
556
-
557
-func (s *RelaySet) ReplaceKnownRelayURLs(relayURLs []string) {
558
- if s == nil {
559
- return
560
- }
561
- s.mu.Lock()
562
- defer s.mu.Unlock()
563
- filtered := make([]string, 0, len(relayURLs))
564
- for _, relayURL := range relayURLs {
565
- relayURL = strings.TrimSpace(relayURL)
566
- if relayURL == "" {
567
- continue
568
- }
569
- duplicate := false
570
- for _, existing := range filtered {
571
- if existing == relayURL {
572
- duplicate = true
573
- break
574
- }
575
- }
576
- if duplicate {
577
- continue
578
- }
579
- filtered = append(filtered, relayURL)
353
+ projections := s.descriptorProjectionsLocked()
354
+ if len(projections) == 0 {
355
+ return nil
356
}
581
- s.knownRelayURLs = append([]string(nil), filtered...)
582
-}
357
584
-func (s *RelaySet) SetBootstrapRelayURLs(relayURLs []string) {
585
- if s == nil {
586
- return
587
- }
588
- normalized := make([]string, 0, len(relayURLs))
589
- seen := make(map[string]struct{}, len(relayURLs))
590
- for _, relayURL := range relayURLs {
591
- relayURL = strings.TrimSpace(relayURL)
592
- if relayURL == "" {
593
- continue
594
- }
595
- parsed, err := utils.NormalizeRelayURL(relayURL)
596
- if err != nil {
597
- log.Warn().
598
- Err(err).
599
- Str("relay_url", relayURL).
600
- Msg("skip invalid bootstrap relay url")
601
- continue
602
- }
603
- if _, ok := seen[parsed]; ok {
358
+ view := make(map[string]RelayState, len(projections))
359
+ for _, projection := range projections {
360
+ relayKey := projection.state.Descriptor.Key()
361
+ if relayKey == "" {
362
continue
363
}
606
- seen[parsed] = struct{}{}
607
- normalized = append(normalized, parsed)
608
- }
609
-
610
- s.mu.Lock()
611
- defer s.mu.Unlock()
612
-
613
- filtered := normalized[:0]
614
- for _, relayURL := range normalized {
615
- if s.isSelfRelayURLLocked(relayURL) {
616
- continue
364
+ expired := relayExpiredAt(projection.state, now)
365
+ view[relayKey] = RelayState{
366
+ Descriptor: projection.state.Descriptor,
367
+ FirstSeenAt: projection.state.FirstSeenAt,
368
+ LastSeenAt: projection.state.LastSeenAt,
369
+ Banned: projection.state.Banned,
370
+ Expired: expired,
371
}
618
- filtered = append(filtered, relayURL)
619
- state := s.localByURL[relayURL]
620
- state.Bootstrap = true
621
- s.localByURL[relayURL] = state
372
}
623
- normalized = filtered
624
-
625
- current := make(map[string]struct{}, len(normalized))
626
- for _, relayURL := range normalized {
627
- current[relayURL] = struct{}{}
628
- }
629
-
630
- existing := make(map[string]struct{}, len(s.knownRelayURLs))
631
- for _, relayURL := range s.knownRelayURLs {
632
- existing[relayURL] = struct{}{}
633
- if _, ok := current[relayURL]; !ok {
634
- state := s.localByURL[relayURL]
635
- if state.Bootstrap {
636
- state.Bootstrap = false
637
- s.localByURL[relayURL] = state
638
- }
639
- }
373
+ if len(view) == 0 {
374
+ return nil
375
}
641
- s.knownRelayURLs = append([]string(nil), normalized...)
642
- s.logStatusChange()
376
+ return view
377
}
378
645
-func (s *RelaySet) mergeKnownRelayURLs(relayURLs []string) error {
646
- if s == nil || len(relayURLs) == 0 {
647
- return nil
648
- }
649
- normalized, err := utils.NormalizeRelayURLs(relayURLs...)
379
+func (s *RelaySet) SetBootstrapRelayURLs(inputs []string) error {
380
+ normalized, err := utils.NormalizeRelayURLs(inputs...)
381
if err != nil {
382
return err
383
}
653
- if len(normalized) == 0 {
654
- return nil
655
- }
384
385
s.mu.Lock()
386
defer s.mu.Unlock()
387
660
- existing := make(map[string]struct{}, len(s.knownRelayURLs))
661
- for _, relayURL := range s.knownRelayURLs {
662
- existing[relayURL] = struct{}{}
388
+ filtered := utils.RemoveRelayURL(normalized, s.selfRelayURL)
389
+
390
+ keep := make(map[string]struct{}, len(filtered))
391
+ for _, relayURL := range filtered {
392
+ keep[relayURL] = struct{}{}
393
}
394
665
- changed := false
666
- for _, relayURL := range normalized {
667
- if s.isSelfRelayURLLocked(relayURL) {
395
+ for _, relayURL := range s.knownRelayURLs {
396
+ if _, ok := keep[relayURL]; ok {
397
continue
398
}
670
- if _, ok := existing[relayURL]; ok {
399
+ if _, ok := s.relayKeysByURL[relayURL]; ok {
400
continue
401
}
673
- existing[relayURL] = struct{}{}
674
- s.knownRelayURLs = append(s.knownRelayURLs, relayURL)
675
- state := s.localByURL[relayURL]
676
- s.localByURL[relayURL] = state
677
- changed = true
678
- }
679
- if changed {
680
- s.logStatusChange()
402
+ if state := s.localByURL[relayURL]; state.isDefaultLocalState() {
403
+ delete(s.localByURL, relayURL)
404
+ }
405
}
406
+
407
+ s.knownRelayURLs = append([]string(nil), filtered...)
408
return nil
409
}
410
411
func (s *RelaySet) registerDescriptor(desc types.RelayDescriptor, now time.Time) (string, bool, bool, error) {
686
- if s == nil {
687
- return "", false, false, nil
688
- }
412
normalized, err := NormalizeDescriptor(desc)
413
if err != nil {
414
return "", false, false, err
@@ -697,131 +420,123 @@ func (s *RelaySet) registerDescriptor(desc types.RelayDescriptor, now time.Time)
420
if knownRelayKey, ok := s.relayKeysByURL[normalized.APIHTTPSAddr]; ok && knownRelayKey != relayKey {
421
return "", false, false, errors.New("descriptor identity does not match known relay url")
422
}
700
-
423
if now.IsZero() {
424
now = time.Now().UTC()
425
}
426
705
- view, ok := s.relays[relayKey]
427
+ record, ok := s.relays[relayKey]
428
added := !ok
429
if !ok {
708
- view.FirstSeenAt = now
430
+ record.FirstSeenAt = now
431
}
710
- previousURL := view.Descriptor.APIHTTPSAddr
711
- previousDescriptor := view.Descriptor
712
- view.Descriptor = normalized
713
- view.LastSeenAt = now
714
- s.relays[relayKey] = view
432
+ previousURL := record.Descriptor.APIHTTPSAddr
433
+ previousDescriptor := record.Descriptor
434
+ record.Descriptor = normalized
435
+ record.LastSeenAt = now
436
+ s.relays[relayKey] = record
437
s.relayKeysByURL[normalized.APIHTTPSAddr] = relayKey
438
if previousURL != "" && previousURL != normalized.APIHTTPSAddr {
439
delete(s.relayKeysByURL, previousURL)
440
+ state := s.localByURL[previousURL]
441
+ bootstrap := false
442
+ if !s.isSelfRelayURLLocked(previousURL) {
443
+ if slices.Contains(s.knownRelayURLs, previousURL) {
444
+ bootstrap = true
445
+ }
446
+ }
447
+ if !bootstrap && state.isDefaultLocalState() {
448
+ delete(s.localByURL, previousURL)
449
+ }
450
}
451
452
changed := added || !reflect.DeepEqual(previousDescriptor, normalized)
453
return relayKey, added, changed, nil
454
}
455
724
-func relayDiscoveryURLs(selfDescriptor types.RelayDescriptor, relayDescriptors []types.RelayDescriptor) []string {
725
- relayURLs := make([]string, 0, 1+len(relayDescriptors))
726
- if apiURL := selfDescriptor.APIHTTPSAddr; apiURL != "" {
727
- relayURLs = append(relayURLs, apiURL)
728
- }
729
- for _, relayDescriptor := range relayDescriptors {
730
- if apiURL := relayDescriptor.APIHTTPSAddr; apiURL != "" {
731
- relayURLs = append(relayURLs, apiURL)
732
- }
733
- }
734
- if len(relayURLs) == 0 {
735
- return nil
736
- }
737
- return relayURLs
738
-}
739
-
740
-func (s *RelaySet) applyDiscoveryDescriptors(targetIdentity types.Identity, targetURL string, selfDescriptor types.RelayDescriptor, relayDescriptors []types.RelayDescriptor, now time.Time) (relaySetChanged bool, addedRelayCount int, err error) {
741
- if s == nil {
742
- return false, 0, nil
743
- }
456
+func (s *RelaySet) applyDiscoveryDescriptorsLocked(targetIdentity types.Identity, targetURL string, selfDescriptor types.RelayDescriptor, relayDescriptors []types.RelayDescriptor, now time.Time) (relaySetChanged bool, err error) {
457
if strings.TrimSpace(targetIdentity.Name) == "" && strings.TrimSpace(targetIdentity.Address) == "" {
745
- return false, 0, errors.New("target relay identity is required")
458
+ return false, errors.New("target relay identity is required")
459
}
460
if now.IsZero() {
461
now = time.Now().UTC()
462
}
463
if err := ValidateDescriptorTarget(selfDescriptor, targetIdentity, targetURL); err != nil {
751
- return false, 0, err
464
+ return false, err
465
}
466
754
- apply := func(desc types.RelayDescriptor, advertise, countAdded bool) error {
467
+ apply := func(desc types.RelayDescriptor, advertise bool) error {
468
+ if !advertise && s.isSelfRelayDescriptorLocked(desc) {
469
+ return nil
470
+ }
471
+
472
_, added, descriptorChanged, err := s.registerDescriptor(desc, now)
473
if err != nil {
474
return err
475
}
476
+
477
localState := s.localByURL[desc.APIHTTPSAddr]
760
- wasAdvertised := localState.Advertised
761
- wasExpired := localState.Expired
478
+ previousState := localState
479
if advertise {
763
- localState.Advertised = true
480
+ localState.status = relayStatusConfirmed
481
+ localState.consecutiveFailures = 0
482
+ } else if localState.status != relayStatusConfirmed {
483
+ localState.status = relayStatusHinted
484
+ localState.consecutiveFailures = 0
485
}
765
- localState.Expired = false
766
- s.localByURL[desc.APIHTTPSAddr] = localState
486
+ s.storeLocalStateLocked(desc.APIHTTPSAddr, localState)
487
768
- changed := added || descriptorChanged || advertise && !wasAdvertised || wasExpired
769
- if added && countAdded {
770
- addedRelayCount++
771
- }
488
+ changed := added || descriptorChanged || !reflect.DeepEqual(previousState, localState)
489
if changed {
490
relaySetChanged = true
491
}
492
return nil
493
}
494
778
- if err := apply(selfDescriptor, true, false); err != nil {
779
- return false, 0, err
495
+ if err := apply(selfDescriptor, true); err != nil {
496
+ return false, err
497
}
498
for _, relayDescriptor := range relayDescriptors {
782
- if s.isSelfRelayDescriptorLocked(relayDescriptor) {
783
- continue
784
- }
785
- if err := apply(relayDescriptor, false, true); err != nil {
786
- return false, 0, err
499
+ if err := apply(relayDescriptor, false); err != nil {
500
+ return false, err
501
}
502
}
503
state := s.localByURL[selfDescriptor.APIHTTPSAddr]
790
- state.Reachable = true
791
- state.ConsecutiveFailures = 0
792
- state.LastSuccessAt = now
793
- s.localByURL[selfDescriptor.APIHTTPSAddr] = state
794
- s.logStatusChange()
795
- return relaySetChanged, addedRelayCount, nil
504
+ state.status = relayStatusConfirmed
505
+ state.consecutiveFailures = 0
506
+ s.storeLocalStateLocked(selfDescriptor.APIHTTPSAddr, state)
507
+ return relaySetChanged, nil
508
}
509
798
-func (s *RelaySet) ApplyRelayDiscoveryResponse(targetIdentity types.Identity, targetURL string, resp types.DiscoveryResponse, now time.Time) (relayURLs []string, relaySetChanged bool, addedRelayCount int, warnErr error, err error) {
510
+func (s *RelaySet) ApplyRelayDiscoveryResponse(targetIdentity types.Identity, targetURL string, resp types.DiscoveryResponse, now time.Time) (relaySetChanged bool, warnErr error, err error) {
511
selfDescriptor, relayDescriptors, validateErr := ValidateRelayDiscoveryResponse(resp, now)
512
warnErr = validateErr
513
if selfDescriptor.Key() == "" {
802
- return nil, false, 0, warnErr, validateErr
514
+ return false, warnErr, validateErr
515
}
516
s.mu.Lock()
805
- relaySetChanged, addedRelayCount, err = s.applyDiscoveryDescriptors(targetIdentity, targetURL, selfDescriptor, relayDescriptors, now)
517
+ relaySetChanged, err = s.applyDiscoveryDescriptorsLocked(targetIdentity, targetURL, selfDescriptor, relayDescriptors, now)
518
s.mu.Unlock()
519
if err != nil {
808
- return nil, false, 0, warnErr, err
520
+ return false, warnErr, err
521
}
810
- return relayDiscoveryURLs(selfDescriptor, relayDescriptors), relaySetChanged, addedRelayCount, warnErr, nil
522
+ return relaySetChanged, warnErr, nil
523
}
524
813
-func (s *RelaySet) ApplyOverlayRelayDiscoveryResponse(targetIdentity types.Identity, targetURL string, resp types.DiscoveryResponse, now time.Time) (relayURLs []string, relaySetChanged bool, addedRelayCount int, warnErr error, err error) {
525
+func (s *RelaySet) ApplyOverlayRelayDiscoveryResponse(targetIdentity types.Identity, targetURL string, resp types.DiscoveryResponse, now time.Time) (relaySetChanged bool, warnErr error, err error) {
526
selfDescriptor, relayDescriptors, validateErr := ValidateRelayDiscoveryResponse(resp, now)
527
warnErr = validateErr
528
if selfDescriptor.Key() == "" {
817
- return nil, false, 0, warnErr, validateErr
529
+ return false, warnErr, validateErr
530
}
531
if err := RequireOverlayRelayDescriptor(selfDescriptor); err != nil {
820
- return nil, false, 0, warnErr, err
532
+ return false, warnErr, err
533
}
534
535
filteredRelayDescriptors := make([]types.RelayDescriptor, 0, len(relayDescriptors))
536
for _, relayDescriptor := range relayDescriptors {
537
+ if s.isSelfRelayDescriptorLocked(relayDescriptor) {
538
+ continue
539
+ }
540
if err := RequireOverlayRelayDescriptor(relayDescriptor); err != nil {
541
if warnErr == nil {
542
warnErr = err
@@ -832,198 +547,43 @@ func (s *RelaySet) ApplyOverlayRelayDiscoveryResponse(targetIdentity types.Ident
547
}
548
549
s.mu.Lock()
835
- relaySetChanged, addedRelayCount, err = s.applyDiscoveryDescriptors(targetIdentity, targetURL, selfDescriptor, filteredRelayDescriptors, now)
550
+ relaySetChanged, err = s.applyDiscoveryDescriptorsLocked(targetIdentity, targetURL, selfDescriptor, filteredRelayDescriptors, now)
551
s.mu.Unlock()
552
if err != nil {
838
- return nil, false, 0, warnErr, err
553
+ return false, warnErr, err
554
}
840
- return relayDiscoveryURLs(selfDescriptor, filteredRelayDescriptors), relaySetChanged, addedRelayCount, warnErr, nil
555
+ return relaySetChanged, warnErr, nil
556
}
557
843
-func (s *RelaySet) ApplyRelayDiscoveryResponseSimple(targetIdentity types.Identity, targetURL string, resp types.DiscoveryResponse, now time.Time) error {
844
- relayURLs, _, _, warnErr, err := s.ApplyRelayDiscoveryResponse(targetIdentity, targetURL, resp, now)
845
- if warnErr != nil {
846
- log.Warn().
847
- Err(warnErr).
848
- Str("relay", targetURL).
849
- Msg("relay discovery response completed with warnings")
850
- }
851
- if err != nil {
852
- return err
853
- }
854
-
855
- targetRelayURL := strings.TrimSpace(resp.Self.APIHTTPSAddr)
856
- if targetRelayURL == "" && len(relayURLs) > 0 {
857
- targetRelayURL = relayURLs[0]
858
- }
859
- if targetRelayURL == "" {
860
- return nil
861
- }
862
- if err := s.mergeKnownRelayURLs([]string{targetRelayURL}); err != nil {
863
- return err
864
- }
865
- return nil
866
-}
867
-
868
-func (s *RelaySet) RegisterBootstrapRelayURLs(inputs []string) ([]string, error) {
869
- if s == nil || len(inputs) == 0 {
870
- return nil, nil
871
- }
872
-
873
- normalized, err := utils.NormalizeRelayURLs(inputs...)
874
- if err != nil {
875
- return nil, err
876
- }
877
- normalized, err = utils.ExcludeLocalRelayURLs(normalized...)
878
- if err != nil {
879
- return nil, err
880
- }
881
- if len(normalized) == 0 {
882
- return nil, nil
883
- }
884
- s.mu.Lock()
885
- defer s.mu.Unlock()
886
-
887
- existing := make(map[string]struct{}, len(s.knownRelayURLs))
888
- for _, relayURL := range s.knownRelayURLs {
889
- existing[relayURL] = struct{}{}
890
- }
891
- added := make([]string, 0, len(normalized))
892
- for _, relayURL := range normalized {
893
- if _, ok := existing[relayURL]; ok {
894
- continue
895
- }
896
- existing[relayURL] = struct{}{}
897
- s.knownRelayURLs = append(s.knownRelayURLs, relayURL)
898
- added = append(added, relayURL)
899
- }
900
- for _, relayURL := range normalized {
901
- state := s.localByURL[relayURL]
902
- state.Bootstrap = true
903
- state.Reachable = false
904
- s.localByURL[relayURL] = state
905
- }
906
- s.logStatusChange()
907
- if len(added) == 0 {
908
- return nil, nil
909
- }
910
- return added, nil
911
-}
912
-
913
-func (s *RelaySet) refreshBootstrapDiscovery(ctx context.Context, rootCAPEM []byte) error {
914
- bootstraps := s.BootstrapDescriptors()
915
- for _, bootstrap := range bootstraps {
916
- resp, err := DiscoverRelayDiscovery(ctx, bootstrap.APIHTTPSAddr, rootCAPEM, nil)
917
- if err != nil {
918
- if ctx.Err() != nil {
919
- return ctx.Err()
920
- }
921
- s.RecordBootstrapDiscoveryFailure(bootstrap.APIHTTPSAddr, err, time.Now().UTC())
922
- continue
923
- }
924
-
925
- now := time.Now().UTC()
926
- relayURLs, _, _, warnErr, applyErr := s.ApplyRelayDiscoveryResponse(bootstrap.Identity, bootstrap.APIHTTPSAddr, resp, now)
927
- if warnErr != nil {
928
- log.Warn().
929
- Err(warnErr).
930
- Str("relay", bootstrap.APIHTTPSAddr).
931
- Msg("bootstrap relay discovery completed with warnings")
932
- }
933
- if applyErr != nil {
934
- s.MarkRelayFailure(bootstrap.APIHTTPSAddr, now)
935
- log.Warn().
936
- Err(applyErr).
937
- Str("relay", bootstrap.APIHTTPSAddr).
938
- Msg("bootstrap relay discovery failed")
939
- continue
940
- }
941
- if len(relayURLs) == 0 {
942
- continue
943
- }
944
- if err := s.mergeKnownRelayURLs(relayURLs); err != nil {
945
- log.Warn().
946
- Err(err).
947
- Str("relay", bootstrap.APIHTTPSAddr).
948
- Msg("merge discovered relay urls")
949
- }
950
- }
951
- return nil
952
-}
953
-
954
-func (s *RelaySet) RunLoop(ctx context.Context, rootCAPEM []byte, syncRuntime func() error) error {
955
- if s == nil {
956
- <-ctx.Done()
957
- return nil
958
- }
959
-
960
- ticker := time.NewTicker(types.DiscoveryPollInterval)
961
- defer ticker.Stop()
962
-
963
- for {
964
- if err := s.refreshBootstrapDiscovery(ctx, rootCAPEM); err != nil {
965
- return err
966
- }
967
- if ctx.Err() != nil {
968
- return nil
969
- }
970
- if syncRuntime != nil {
971
- if err := syncRuntime(); err != nil {
972
- return err
973
- }
974
- }
975
-
976
- select {
977
- case <-ctx.Done():
978
- return nil
979
- case <-ticker.C:
980
- }
981
- }
982
-}
983
-
984
-func (s *RelaySet) RecordDiscoveryFailure(identity types.Identity, relayURL string, err error) (expired bool, expireReason string, consecutiveFailures int) {
985
- return s.RecordDiscoveryFailureWithRecovery(identity, relayURL, err, defaultDiscoveryRecoveryFailures, time.Now().UTC())
986
-}
987
-
988
-func (s *RelaySet) RecordDiscoveryFailureWithRecovery(identity types.Identity, relayURL string, err error, recoveryFailures int, now time.Time) (expired bool, expireReason string, consecutiveFailures int) {
989
- if s == nil {
990
- return false, "", 0
991
- }
558
+func (s *RelaySet) RecordDiscoveryFailure(identity types.Identity, relayURL string, err error, recoveryFailures int) (expired bool, expireReason string, consecutiveFailures int) {
559
relayKey := identity.Key()
560
if relayKey == "" {
561
return false, "", 0
562
}
563
relayURL = strings.TrimSpace(relayURL)
997
- if relayURL == "" {
998
- return false, "", 0
999
- }
1000
- if now.IsZero() {
1001
- now = time.Now().UTC()
1002
- } else {
1003
- now = now.UTC()
1004
- }
564
565
s.mu.Lock()
566
defer s.mu.Unlock()
567
1009
- view, ok := s.relays[relayKey]
568
+ record, ok := s.relays[relayKey]
569
if !ok {
570
return false, "", 0
571
}
572
+ if relayURL == "" || s.relayKeysByURL[relayURL] != relayKey {
573
+ relayURL = record.Descriptor.APIHTTPSAddr
574
+ }
575
+ if relayURL == "" {
576
+ return false, "", 0
577
+ }
578
579
localState := s.localByURL[relayURL]
1015
- localState.Reachable = false
1016
- localState.ConsecutiveFailures++
1017
- localState.LastFailureAt = now
1018
- s.localByURL[relayURL] = localState
1019
- s.logStatusChange()
1020
- if !localState.Expired && localState.ConsecutiveFailures >= recoveryFailures {
1021
- state := s.localByURL[view.Descriptor.APIHTTPSAddr]
1022
- state.Advertised = false
1023
- state.Expired = true
1024
- s.localByURL[view.Descriptor.APIHTTPSAddr] = state
1025
- s.logStatusChange()
1026
- return true, "recovery", localState.ConsecutiveFailures
580
+ localState.consecutiveFailures++
581
+ s.storeLocalStateLocked(relayURL, localState)
582
+ if localState.status != relayStatusExpired && localState.consecutiveFailures >= recoveryFailures {
583
+ state := s.localByURL[record.Descriptor.APIHTTPSAddr]
584
+ state.status = relayStatusExpired
585
+ s.storeLocalStateLocked(record.Descriptor.APIHTTPSAddr, state)
586
+ return true, "recovery", localState.consecutiveFailures
587
}
588
589
var apiErr *types.APIRequestError
@@ -1031,12 +591,10 @@ func (s *RelaySet) RecordDiscoveryFailureWithRecovery(identity types.Identity, r
591
(apiErr.StatusCode == http.StatusForbidden ||
592
apiErr.StatusCode == http.StatusNotFound ||
593
apiErr.StatusCode == http.StatusGone) {
1034
- state := s.localByURL[view.Descriptor.APIHTTPSAddr]
1035
- state.Advertised = false
1036
- state.Expired = true
1037
- s.localByURL[view.Descriptor.APIHTTPSAddr] = state
1038
- s.logStatusChange()
1039
- return true, "status", localState.ConsecutiveFailures
594
+ state := s.localByURL[record.Descriptor.APIHTTPSAddr]
595
+ state.status = relayStatusExpired
596
+ s.storeLocalStateLocked(record.Descriptor.APIHTTPSAddr, state)
597
+ return true, "status", localState.consecutiveFailures
598
}
1041
- return false, "", localState.ConsecutiveFailures
599
+ return false, "", localState.consecutiveFailures
600
}
portal/server.go
+26
-120
@@ -29,13 +29,12 @@ import (
29
)
30
31
const (
32
- defaultLeaseTTL = 30 * time.Second
33
- defaultClaimTimeout = 10 * time.Second
34
- defaultIdleKeepalive = 15 * time.Second
35
- defaultReadyQueueLimit = 8
36
- defaultClientHelloWait = 2 * time.Second
37
- defaultControlBodyLimit = 4 << 20
38
- defaultWGRecoveryFailures = 3
32
+ defaultLeaseTTL = 30 * time.Second
33
+ defaultClaimTimeout = 10 * time.Second
34
+ defaultIdleKeepalive = 15 * time.Second
35
+ defaultReadyQueueLimit = 8
36
+ defaultClientHelloWait = 2 * time.Second
37
+ defaultControlBodyLimit = 4 << 20
38
)
39
40
type ServerConfig struct {
@@ -81,7 +80,6 @@ type Server struct {
80
identity types.Identity
81
cfg ServerConfig
82
trustedProxyCIDRs []*net.IPNet
84
- wgConfig wireguard.Config
83
relaySet *discovery.RelaySet
84
thumbnails *thumbnail.Service
85
shutdownOnce sync.Once
@@ -213,7 +211,6 @@ func NewServer(cfg ServerConfig) (*Server, error) {
211
loadMgr: policy.NewLoadManager(),
212
identity: identity,
213
trustedProxyCIDRs: trustedProxyCIDRs,
216
- wgConfig: wgConfig,
214
thumbnails: thumbnail.NewService(cfg.HeadlessShellURL),
215
}
216
if cfg.DiscoveryEnabled {
@@ -221,7 +218,7 @@ func NewServer(cfg ServerConfig) (*Server, error) {
218
if err := set.SetSelfRelay(identity, selfRelayURL); err != nil {
219
return nil, fmt.Errorf("set self relay: %w", err)
220
}
224
- if _, err := set.RegisterBootstrapRelayURLs(cfg.Bootstraps); err != nil {
221
+ if err := set.SetBootstrapRelayURLs(cfg.Bootstraps); err != nil {
222
return nil, err
223
}
224
s.relaySet = set
@@ -273,7 +270,7 @@ func (s *Server) Start(ctx context.Context, apiMux *http.ServeMux) error {
270
s.cancel = cancel
271
s.group = group
272
276
- if s.relaySet != nil && strings.TrimSpace(s.wgConfig.PrivateKey) != "" {
273
+ if s.relaySet != nil && strings.TrimSpace(s.cfg.WireGuardPrivateKey) != "" {
274
if err := s.startOverlay(); err != nil {
275
acmeManager.Stop()
276
_ = apiServer.Close()
@@ -417,7 +414,7 @@ func (s *Server) wireGuardOverlayEnabled() bool {
414
if s == nil {
415
return false
416
}
420
- return strings.TrimSpace(s.wgConfig.PrivateKey) != ""
417
+ return strings.TrimSpace(s.cfg.WireGuardPrivateKey) != ""
418
}
419
420
func (s *Server) LeaseSnapshots() []types.Lease {
@@ -672,12 +669,19 @@ func (s *Server) startOverlay() error {
669
peerMux.HandleFunc(types.PathDiscovery, s.handleRelayDiscovery)
670
}
671
675
- overlay, err := wireguard.NewOverlay(s.wgConfig, peerMux)
672
+ overlay, err := wireguard.NewOverlay(wireguard.Config{
673
+ PrivateKey: s.cfg.WireGuardPrivateKey,
674
+ PublicKey: s.cfg.WireGuardPublicKey,
675
+ Endpoint: s.cfg.WireGuardEndpoint,
676
+ OverlayIPv4: s.cfg.OverlayIPv4,
677
+ OverlayCIDRs: s.cfg.OverlayCIDRs,
678
+ ListenPort: s.cfg.DiscoveryPort,
679
+ }, peerMux)
680
if err != nil {
681
return fmt.Errorf("start wireguard overlay: %w", err)
682
}
683
680
- if err := overlay.Sync(s.identity.Key(), s.relaySet.Snapshot()); err != nil {
684
+ if err := overlay.Sync(s.relaySet.View()); err != nil {
685
_ = overlay.Shutdown(context.Background())
686
return fmt.Errorf("sync wireguard peers: %w", err)
687
}
@@ -691,117 +695,19 @@ func (s *Server) runRelayDiscoveryLoop(ctx context.Context) error {
695
<-ctx.Done()
696
return nil
697
}
698
+ refresher, err := discovery.NewRefresher(s.relaySet, nil, s.overlay)
699
+ if err != nil {
700
+ return err
701
+ }
702
ticker := time.NewTicker(types.DiscoveryPollInterval)
703
defer ticker.Stop()
704
705
for {
698
- bootstraps := s.relaySet.BootstrapDescriptors()
699
-
700
- for _, bootstrap := range bootstraps {
701
- resp, err := discovery.DiscoverRelayDiscovery(ctx, bootstrap.APIHTTPSAddr, nil, nil)
702
- if err != nil {
703
- if ctx.Err() != nil {
704
- return nil
705
- }
706
- s.relaySet.RecordBootstrapDiscoveryFailure(bootstrap.APIHTTPSAddr, err, time.Now().UTC())
707
- continue
708
- }
709
-
710
- now := time.Now().UTC()
711
- var relaySetChanged bool
712
- var warnErr error
713
- _, relaySetChanged, _, warnErr, err = s.relaySet.ApplyRelayDiscoveryResponse(bootstrap.Identity, bootstrap.APIHTTPSAddr, resp, now)
714
- if relaySetChanged && s.overlay != nil {
715
- if syncErr := s.overlay.Sync(s.identity.Key(), s.relaySet.Snapshot()); syncErr != nil {
716
- if warnErr == nil {
717
- warnErr = syncErr
718
- }
719
- }
720
- }
721
- if err != nil {
722
- s.relaySet.MarkRelayFailure(bootstrap.APIHTTPSAddr, time.Now().UTC())
723
- log.Warn().
724
- Err(err).
725
- Str("relay", bootstrap.APIHTTPSAddr).
726
- Msg("bootstrap relay discovery failed")
727
- continue
728
- }
729
- if warnErr != nil {
730
- log.Warn().
731
- Err(warnErr).
732
- Str("relay", bootstrap.APIHTTPSAddr).
733
- Msg("bootstrap relay discovery completed with warnings")
734
- }
735
- }
736
- if ctx.Err() != nil {
737
- return nil
738
- }
739
-
740
- if s.overlay != nil {
741
- overlayClient := s.overlay.Client()
742
- syncableRelays := s.relaySet.SyncableDescriptors()
743
-
744
- for _, relay := range syncableRelays {
745
- var failureErr error
746
-
747
- if err := discovery.RequireOverlayRelayDescriptor(relay); err != nil {
748
- failureErr = err
749
- } else {
750
- discoverURL := "http://" + net.JoinHostPort(relay.OverlayIPv4, fmt.Sprintf("%d", wireguard.DefaultPeerAPIHTTPPort))
751
- resp, err := discovery.DiscoverRelayDiscovery(ctx, discoverURL, nil, overlayClient)
752
- if err != nil {
753
- if ctx.Err() != nil {
754
- return nil
755
- }
756
- failureErr = err
757
- } else {
758
- now := time.Now().UTC()
759
- var relaySetChanged bool
760
- var warnErr error
761
- var snapshot map[string]types.RelayState
762
- _, relaySetChanged, _, warnErr, err = s.relaySet.ApplyOverlayRelayDiscoveryResponse(relay.Identity, relay.APIHTTPSAddr, resp, now)
763
- if relaySetChanged {
764
- snapshot = s.relaySet.Snapshot()
765
- if syncErr := s.overlay.Sync(s.identity.Key(), snapshot); syncErr != nil {
766
- if warnErr == nil {
767
- warnErr = syncErr
768
- }
769
- }
770
- }
771
- if err != nil {
772
- failureErr = err
773
- } else {
774
- if warnErr != nil {
775
- log.Warn().
776
- Err(warnErr).
777
- Str("relay", relay.APIHTTPSAddr).
778
- Msg("overlay relay discovery completed with warnings")
779
- }
780
- continue
781
- }
782
- }
783
- }
784
-
785
- expired, expireReason, consecutiveFailures := s.relaySet.RecordDiscoveryFailureWithRecovery(relay.Identity, relay.APIHTTPSAddr, failureErr, defaultWGRecoveryFailures, time.Now().UTC())
786
- if expired {
787
- if syncErr := s.overlay.Sync(s.identity.Key(), s.relaySet.Snapshot()); syncErr != nil && failureErr == nil {
788
- failureErr = syncErr
789
- }
790
- }
791
-
792
- event := log.Warn().
793
- Err(failureErr).
794
- Str("relay", relay.APIHTTPSAddr)
795
- if expired {
796
- event = event.
797
- Bool("expired", true).
798
- Str("reason", expireReason)
799
- if consecutiveFailures > 0 {
800
- event = event.Int("consecutive_failures", consecutiveFailures)
801
- }
802
- }
803
- event.Msg("overlay relay discovery failed")
706
+ if err := refresher.Refresh(ctx); err != nil {
707
+ if ctx.Err() != nil {
708
+ return nil
709
}
710
+ return err
711
}
712
if ctx.Err() != nil {
713
return nil
portal/server_test.go
+107
-25
@@ -51,7 +51,11 @@ func mustRelayDescriptor(t *testing.T, relayURL string) types.RelayDescriptor {
51
52
func applyRelay(t *testing.T, set *discovery.RelaySet, identity types.Identity, targetURL string, resp types.DiscoveryResponse, now time.Time) error {
53
t.Helper()
54
- return set.ApplyRelayDiscoveryResponseSimple(identity, targetURL, resp, now)
54
+ _, warnErr, err := set.ApplyRelayDiscoveryResponse(identity, targetURL, resp, now)
55
+ if err != nil {
56
+ return err
57
+ }
58
+ return warnErr
59
}
60
61
func tempIdentityPath(t *testing.T) string {
@@ -498,13 +502,15 @@ func TestServerSetBootstrapRelayURLsAllowsLoopbackButSkipsSelfRelay(t *testing.T
502
t.Fatalf("NewServer() error = %v", err)
503
}
504
501
- server.relaySet.SetBootstrapRelayURLs([]string{
505
+ if err := server.relaySet.SetBootstrapRelayURLs([]string{
506
"https://bootstrap.example.com",
507
"https://localhost:4017",
508
"https://relay-a.example.com",
509
"https://relay-b.example.com",
506
- })
507
- advertisedDescriptors := server.relaySet.ActiveRelayDescriptors()
510
+ }); err != nil {
511
+ t.Fatalf("SetBootstrapRelayURLs() error = %v", err)
512
+ }
513
+ advertisedDescriptors := server.relaySet.AdvertisedDescriptors()
514
knownURLs := append([]string(nil), server.relaySet.ActiveRelayURLs()...)
515
sort.Strings(knownURLs)
516
if !reflect.DeepEqual(knownURLs, []string{
@@ -608,7 +614,7 @@ func TestServerRecordVerifiedDiscoveryPeerRequiresDirectConfirmation(t *testing.
614
if !reflect.DeepEqual(knownURLs, []string{"https://bootstrap.example.com"}) {
615
t.Fatalf("ActiveRelayURLs() = %v, want [%q]", knownURLs, "https://bootstrap.example.com")
616
}
611
- advertisedDescriptors := server.relaySet.ActiveRelayDescriptors()
617
+ advertisedDescriptors := server.relaySet.AdvertisedDescriptors()
618
advertisedURLs := make([]string, 0, len(advertisedDescriptors))
619
for _, descriptor := range advertisedDescriptors {
620
if strings.TrimSpace(descriptor.APIHTTPSAddr) == "" {
@@ -618,7 +624,7 @@ func TestServerRecordVerifiedDiscoveryPeerRequiresDirectConfirmation(t *testing.
624
}
625
sort.Strings(advertisedURLs)
626
if !reflect.DeepEqual(advertisedURLs, []string{"https://bootstrap.example.com"}) {
621
- t.Fatalf("ActiveRelayDescriptors() = %v, want [%q]", advertisedURLs, "https://bootstrap.example.com")
627
+ t.Fatalf("AdvertisedDescriptors() = %v, want [%q]", advertisedURLs, "https://bootstrap.example.com")
628
}
629
630
err = applyDiscovery(
@@ -629,7 +635,7 @@ func TestServerRecordVerifiedDiscoveryPeerRequiresDirectConfirmation(t *testing.
635
if err != nil {
636
t.Fatalf("applyRelayDiscoveryResponse() confirm error = %v", err)
637
}
632
- advertisedDescriptors = server.relaySet.ActiveRelayDescriptors()
638
+ advertisedDescriptors = server.relaySet.AdvertisedDescriptors()
639
advertisedURLs = advertisedURLs[:0]
640
for _, descriptor := range advertisedDescriptors {
641
if strings.TrimSpace(descriptor.APIHTTPSAddr) == "" {
@@ -639,7 +645,61 @@ func TestServerRecordVerifiedDiscoveryPeerRequiresDirectConfirmation(t *testing.
645
}
646
sort.Strings(advertisedURLs)
647
if !reflect.DeepEqual(advertisedURLs, []string{"https://bootstrap.example.com", "https://relay-a.example.com"}) {
642
- t.Fatalf("ActiveRelayDescriptors() = %v, want [%q %q]", advertisedURLs, "https://bootstrap.example.com", "https://relay-a.example.com")
648
+ t.Fatalf("AdvertisedDescriptors() = %v, want [%q %q]", advertisedURLs, "https://bootstrap.example.com", "https://relay-a.example.com")
649
+ }
650
+}
651
+
652
+func TestServerBannedDiscoveryPeerIsNotAdvertised(t *testing.T) {
653
+ t.Parallel()
654
+
655
+ server, err := NewServer(ServerConfig{
656
+ PortalURL: "https://portal.example.com",
657
+ IdentityPath: tempIdentityPath(t),
658
+ Bootstraps: []string{"https://bootstrap.example.com"},
659
+ DiscoveryEnabled: true,
660
+ })
661
+ if err != nil {
662
+ t.Fatalf("NewServer() error = %v", err)
663
+ }
664
+
665
+ bootstrapDesc := mustRelayDescriptor(t, "https://bootstrap.example.com")
666
+ relayADesc := mustRelayDescriptor(t, "https://relay-a.example.com")
667
+ now := time.Now().UTC()
668
+
669
+ if err := applyRelay(
670
+ t,
671
+ server.relaySet,
672
+ bootstrapDesc.Identity,
673
+ bootstrapDesc.APIHTTPSAddr,
674
+ types.DiscoveryResponse{ProtocolVersion: types.ProtocolVersion, Self: bootstrapDesc, Relays: []types.RelayDescriptor{relayADesc}},
675
+ now,
676
+ ); err != nil {
677
+ t.Fatalf("ApplyRelayDiscoveryResponse() bootstrap error = %v", err)
678
+ }
679
+ if err := applyRelay(
680
+ t,
681
+ server.relaySet,
682
+ relayADesc.Identity,
683
+ relayADesc.APIHTTPSAddr,
684
+ types.DiscoveryResponse{ProtocolVersion: types.ProtocolVersion, Self: relayADesc},
685
+ now.Add(time.Second),
686
+ ); err != nil {
687
+ t.Fatalf("ApplyRelayDiscoveryResponse() direct confirm error = %v", err)
688
+ }
689
+
690
+ server.relaySet.BanRelayURL(relayADesc.APIHTTPSAddr)
691
+
692
+ advertisedDescriptors := server.relaySet.AdvertisedDescriptors()
693
+ advertisedURLs := make([]string, 0, len(advertisedDescriptors))
694
+ for _, descriptor := range advertisedDescriptors {
695
+ if strings.TrimSpace(descriptor.APIHTTPSAddr) == "" {
696
+ continue
697
+ }
698
+ advertisedURLs = append(advertisedURLs, descriptor.APIHTTPSAddr)
699
+ }
700
+ sort.Strings(advertisedURLs)
701
+ if !reflect.DeepEqual(advertisedURLs, []string{"https://bootstrap.example.com"}) {
702
+ t.Fatalf("AdvertisedDescriptors() = %v, want banned relay excluded", advertisedURLs)
703
}
704
}
705
@@ -660,7 +720,9 @@ func TestServerRecordVerifiedDiscoveryPeerExpiresAfterRepeatedDirectFailures(t *
720
relayADesc := mustRelayDescriptor(t, "https://relay-a.example.com")
721
now := time.Now().UTC()
722
663
- if err := server.relaySet.ApplyRelayDiscoveryResponseSimple(
723
+ if err := applyRelay(
724
+ t,
725
+ server.relaySet,
726
bootstrapDesc.Identity,
727
bootstrapDesc.APIHTTPSAddr,
728
types.DiscoveryResponse{ProtocolVersion: types.ProtocolVersion, Self: bootstrapDesc, Relays: []types.RelayDescriptor{relayADesc}},
@@ -668,7 +730,9 @@ func TestServerRecordVerifiedDiscoveryPeerExpiresAfterRepeatedDirectFailures(t *
730
); err != nil {
731
t.Fatalf("ApplyRelayDiscoveryResponse() bootstrap error = %v", err)
732
}
671
- if err := server.relaySet.ApplyRelayDiscoveryResponseSimple(
733
+ if err := applyRelay(
734
+ t,
735
+ server.relaySet,
736
relayADesc.Identity,
737
relayADesc.APIHTTPSAddr,
738
types.DiscoveryResponse{ProtocolVersion: types.ProtocolVersion, Self: relayADesc},
@@ -682,6 +746,7 @@ func TestServerRecordVerifiedDiscoveryPeerExpiresAfterRepeatedDirectFailures(t *
746
relayADesc.Identity,
747
relayADesc.APIHTTPSAddr,
748
errors.New("direct discovery failed"),
749
+ 3,
750
)
751
if consecutiveFailures != attempt {
752
t.Fatalf("RecordDiscoveryFailure() consecutive = %d, want %d", consecutiveFailures, attempt)
@@ -694,7 +759,7 @@ func TestServerRecordVerifiedDiscoveryPeerExpiresAfterRepeatedDirectFailures(t *
759
}
760
}
761
697
- advertisedDescriptors := server.relaySet.ActiveRelayDescriptors()
762
+ advertisedDescriptors := server.relaySet.AdvertisedDescriptors()
763
advertisedURLs := make([]string, 0, len(advertisedDescriptors))
764
for _, descriptor := range advertisedDescriptors {
765
if strings.TrimSpace(descriptor.APIHTTPSAddr) == "" {
@@ -704,7 +769,7 @@ func TestServerRecordVerifiedDiscoveryPeerExpiresAfterRepeatedDirectFailures(t *
769
}
770
sort.Strings(advertisedURLs)
771
if !reflect.DeepEqual(advertisedURLs, []string{"https://bootstrap.example.com"}) {
707
- t.Fatalf("ActiveRelayDescriptors() = %v, want [%q] after relay expiry", advertisedURLs, "https://bootstrap.example.com")
772
+ t.Fatalf("AdvertisedDescriptors() = %v, want [%q] after relay expiry", advertisedURLs, "https://bootstrap.example.com")
773
}
774
}
775
@@ -725,7 +790,9 @@ func TestServerBootstrapHintDoesNotResetDirectFailureBudget(t *testing.T) {
790
relayADesc := mustRelayDescriptor(t, "https://relay-a.example.com")
791
now := time.Now().UTC()
792
728
- if err := server.relaySet.ApplyRelayDiscoveryResponseSimple(
793
+ if err := applyRelay(
794
+ t,
795
+ server.relaySet,
796
bootstrapDesc.Identity,
797
bootstrapDesc.APIHTTPSAddr,
798
types.DiscoveryResponse{ProtocolVersion: types.ProtocolVersion, Self: bootstrapDesc, Relays: []types.RelayDescriptor{relayADesc}},
@@ -733,7 +800,9 @@ func TestServerBootstrapHintDoesNotResetDirectFailureBudget(t *testing.T) {
800
); err != nil {
801
t.Fatalf("ApplyRelayDiscoveryResponse() bootstrap error = %v", err)
802
}
736
- if err := server.relaySet.ApplyRelayDiscoveryResponseSimple(
803
+ if err := applyRelay(
804
+ t,
805
+ server.relaySet,
806
relayADesc.Identity,
807
relayADesc.APIHTTPSAddr,
808
types.DiscoveryResponse{ProtocolVersion: types.ProtocolVersion, Self: relayADesc},
@@ -747,6 +816,7 @@ func TestServerBootstrapHintDoesNotResetDirectFailureBudget(t *testing.T) {
816
relayADesc.Identity,
817
relayADesc.APIHTTPSAddr,
818
errors.New("direct discovery failed"),
819
+ 3,
820
)
821
if expired {
822
t.Fatalf("RecordDiscoveryFailure() expired early on attempt %d", attempt)
@@ -755,7 +825,9 @@ func TestServerBootstrapHintDoesNotResetDirectFailureBudget(t *testing.T) {
825
t.Fatalf("RecordDiscoveryFailure() consecutive = %d, want %d", consecutiveFailures, attempt)
826
}
827
758
- if err := server.relaySet.ApplyRelayDiscoveryResponseSimple(
828
+ if err := applyRelay(
829
+ t,
830
+ server.relaySet,
831
bootstrapDesc.Identity,
832
bootstrapDesc.APIHTTPSAddr,
833
types.DiscoveryResponse{ProtocolVersion: types.ProtocolVersion, Self: bootstrapDesc, Relays: []types.RelayDescriptor{relayADesc}},
@@ -764,7 +836,7 @@ func TestServerBootstrapHintDoesNotResetDirectFailureBudget(t *testing.T) {
836
t.Fatalf("ApplyRelayDiscoveryResponse() hinted refresh error = %v", err)
837
}
838
767
- advertisedDescriptors := server.relaySet.ActiveRelayDescriptors()
839
+ advertisedDescriptors := server.relaySet.AdvertisedDescriptors()
840
advertisedURLs := make([]string, 0, len(advertisedDescriptors))
841
for _, descriptor := range advertisedDescriptors {
842
if strings.TrimSpace(descriptor.APIHTTPSAddr) == "" {
@@ -774,7 +846,7 @@ func TestServerBootstrapHintDoesNotResetDirectFailureBudget(t *testing.T) {
846
}
847
sort.Strings(advertisedURLs)
848
if !reflect.DeepEqual(advertisedURLs, []string{"https://bootstrap.example.com", "https://relay-a.example.com"}) {
777
- t.Fatalf("ActiveRelayDescriptors() = %v, want relay to remain advertised before expiry", advertisedURLs)
849
+ t.Fatalf("AdvertisedDescriptors() = %v, want relay to remain advertised before expiry", advertisedURLs)
850
}
851
}
852
@@ -782,6 +854,7 @@ func TestServerBootstrapHintDoesNotResetDirectFailureBudget(t *testing.T) {
854
relayADesc.Identity,
855
relayADesc.APIHTTPSAddr,
856
errors.New("direct discovery failed"),
857
+ 3,
858
)
859
if !expired {
860
t.Fatal("RecordDiscoveryFailure() expired = false on final attempt, want true")
@@ -808,7 +881,9 @@ func TestServerExpiredDiscoveryPeerNeedsFreshDirectConfirmation(t *testing.T) {
881
relayADesc := mustRelayDescriptor(t, "https://relay-a.example.com")
882
now := time.Now().UTC()
883
811
- if err := server.relaySet.ApplyRelayDiscoveryResponseSimple(
884
+ if err := applyRelay(
885
+ t,
886
+ server.relaySet,
887
bootstrapDesc.Identity,
888
bootstrapDesc.APIHTTPSAddr,
889
types.DiscoveryResponse{ProtocolVersion: types.ProtocolVersion, Self: bootstrapDesc, Relays: []types.RelayDescriptor{relayADesc}},
@@ -816,7 +891,9 @@ func TestServerExpiredDiscoveryPeerNeedsFreshDirectConfirmation(t *testing.T) {
891
); err != nil {
892
t.Fatalf("ApplyRelayDiscoveryResponse() bootstrap error = %v", err)
893
}
819
- if err := server.relaySet.ApplyRelayDiscoveryResponseSimple(
894
+ if err := applyRelay(
895
+ t,
896
+ server.relaySet,
897
relayADesc.Identity,
898
relayADesc.APIHTTPSAddr,
899
types.DiscoveryResponse{ProtocolVersion: types.ProtocolVersion, Self: relayADesc},
@@ -830,10 +907,13 @@ func TestServerExpiredDiscoveryPeerNeedsFreshDirectConfirmation(t *testing.T) {
907
relayADesc.Identity,
908
relayADesc.APIHTTPSAddr,
909
errors.New("direct discovery failed"),
910
+ 3,
911
)
912
}
913
836
- if err := server.relaySet.ApplyRelayDiscoveryResponseSimple(
914
+ if err := applyRelay(
915
+ t,
916
+ server.relaySet,
917
bootstrapDesc.Identity,
918
bootstrapDesc.APIHTTPSAddr,
919
types.DiscoveryResponse{ProtocolVersion: types.ProtocolVersion, Self: bootstrapDesc, Relays: []types.RelayDescriptor{relayADesc}},
@@ -842,7 +922,7 @@ func TestServerExpiredDiscoveryPeerNeedsFreshDirectConfirmation(t *testing.T) {
922
t.Fatalf("ApplyRelayDiscoveryResponse() fresh bootstrap error = %v", err)
923
}
924
845
- advertisedDescriptors := server.relaySet.ActiveRelayDescriptors()
925
+ advertisedDescriptors := server.relaySet.AdvertisedDescriptors()
926
advertisedURLs := make([]string, 0, len(advertisedDescriptors))
927
for _, descriptor := range advertisedDescriptors {
928
if strings.TrimSpace(descriptor.APIHTTPSAddr) == "" {
@@ -852,10 +932,12 @@ func TestServerExpiredDiscoveryPeerNeedsFreshDirectConfirmation(t *testing.T) {
932
}
933
sort.Strings(advertisedURLs)
934
if !reflect.DeepEqual(advertisedURLs, []string{"https://bootstrap.example.com"}) {
855
- t.Fatalf("ActiveRelayDescriptors() = %v, want relay to stay hidden until reconfirmed", advertisedURLs)
935
+ t.Fatalf("AdvertisedDescriptors() = %v, want relay to stay hidden until reconfirmed", advertisedURLs)
936
}
937
858
- if err := server.relaySet.ApplyRelayDiscoveryResponseSimple(
938
+ if err := applyRelay(
939
+ t,
940
+ server.relaySet,
941
relayADesc.Identity,
942
relayADesc.APIHTTPSAddr,
943
types.DiscoveryResponse{ProtocolVersion: types.ProtocolVersion, Self: relayADesc},
@@ -864,7 +946,7 @@ func TestServerExpiredDiscoveryPeerNeedsFreshDirectConfirmation(t *testing.T) {
946
t.Fatalf("ApplyRelayDiscoveryResponse() reconfirm error = %v", err)
947
}
948
867
- advertisedDescriptors = server.relaySet.ActiveRelayDescriptors()
949
+ advertisedDescriptors = server.relaySet.AdvertisedDescriptors()
950
advertisedURLs = advertisedURLs[:0]
951
for _, descriptor := range advertisedDescriptors {
952
if strings.TrimSpace(descriptor.APIHTTPSAddr) == "" {
@@ -874,7 +956,7 @@ func TestServerExpiredDiscoveryPeerNeedsFreshDirectConfirmation(t *testing.T) {
956
}
957
sort.Strings(advertisedURLs)
958
if !reflect.DeepEqual(advertisedURLs, []string{"https://bootstrap.example.com", "https://relay-a.example.com"}) {
877
- t.Fatalf("ActiveRelayDescriptors() = %v, want relay restored after direct confirmation", advertisedURLs)
959
+ t.Fatalf("AdvertisedDescriptors() = %v, want relay restored after direct confirmation", advertisedURLs)
960
}
961
}
962
portal/wireguard/overlay.go
+53
-19
@@ -6,14 +6,22 @@ import (
6
"fmt"
7
"net"
8
"net/http"
9
+ "net/url"
10
"sort"
11
"strings"
12
"time"
13
14
+ "github.com/gosuda/portal-tunnel/v2/portal/discovery"
15
"github.com/gosuda/portal-tunnel/v2/types"
16
"github.com/gosuda/portal-tunnel/v2/utils"
17
)
18
19
+type desiredPeer struct {
20
+ wireGuardPublicKey string
21
+ wireGuardEndpoint string
22
+ allowedIPs []string
23
+}
24
+
25
type Config struct {
26
PrivateKey string
27
PublicKey string
@@ -77,12 +85,18 @@ func NormalizeConfig(rootHost string, cfg Config) (Config, error) {
85
}
86
87
type Overlay struct {
80
- stack *stack
81
- listener net.Listener
82
- server *http.Server
88
+ selfWireGuardPublicKey string
89
+ stack *stack
90
+ listener net.Listener
91
+ server *http.Server
92
}
93
94
func NewOverlay(cfg Config, handler http.Handler) (*Overlay, error) {
95
+ selfWireGuardPublicKey := strings.TrimSpace(cfg.PublicKey)
96
+ if selfWireGuardPublicKey == "" {
97
+ return nil, errors.New("wireguard public key is required")
98
+ }
99
+
100
stack, err := newStack(cfg)
101
if err != nil {
102
return nil, err
@@ -100,9 +114,10 @@ func NewOverlay(cfg Config, handler http.Handler) (*Overlay, error) {
114
}
115
116
return &Overlay{
103
- stack: stack,
104
- listener: listener,
105
- server: server,
117
+ selfWireGuardPublicKey: selfWireGuardPublicKey,
118
+ stack: stack,
119
+ listener: listener,
120
+ server: server,
121
}, nil
122
}
123
@@ -154,21 +169,40 @@ func (o *Overlay) Client() *http.Client {
169
}
170
}
171
157
-func (o *Overlay) Sync(selfIdentityKey string, snapshot map[string]types.RelayState) error {
172
+func (o *Overlay) DiscoverRelay(ctx context.Context, relay types.RelayDescriptor) (types.DiscoveryResponse, error) {
173
+ if o == nil || o.stack == nil {
174
+ return types.DiscoveryResponse{}, errors.New("overlay is not initialized")
175
+ }
176
+ if strings.TrimSpace(relay.OverlayIPv4) == "" {
177
+ return types.DiscoveryResponse{}, errors.New("relay overlay ipv4 is required")
178
+ }
179
+
180
+ var resp types.DiscoveryResponse
181
+ baseURL := &url.URL{
182
+ Scheme: "http",
183
+ Host: net.JoinHostPort(relay.OverlayIPv4, fmt.Sprintf("%d", DefaultPeerAPIHTTPPort)),
184
+ }
185
+ if err := utils.HTTPDoAPIPath(ctx, o.Client(), baseURL, http.MethodGet, types.PathDiscovery, nil, nil, &resp); err != nil {
186
+ return types.DiscoveryResponse{}, err
187
+ }
188
+ return resp, nil
189
+}
190
+
191
+func (o *Overlay) Sync(view map[string]discovery.RelayState) error {
192
if o == nil || o.stack == nil {
193
return nil
194
}
161
- return o.stack.ApplyPeers(peersForSnapshot(selfIdentityKey, snapshot))
195
+ return o.stack.ApplyPeers(peersForView(o.selfWireGuardPublicKey, view))
196
}
197
164
-func peersForSnapshot(selfIdentityKey string, snapshot map[string]types.RelayState) []types.DesiredPeer {
165
- peers := make([]types.DesiredPeer, 0, len(snapshot))
166
- for _, state := range snapshot {
167
- if state.Expired {
198
+func peersForView(selfWireGuardPublicKey string, view map[string]discovery.RelayState) []desiredPeer {
199
+ peers := make([]desiredPeer, 0, len(view))
200
+ for _, relay := range view {
201
+ if relay.Expired || relay.Banned {
202
continue
203
}
170
- desc := state.Descriptor
171
- if desc.Key() == selfIdentityKey || !desc.SupportsOverlayPeer {
204
+ desc := relay.Descriptor
205
+ if desc.WireGuardPublicKey == selfWireGuardPublicKey || !desc.SupportsOverlayPeer {
206
continue
207
}
208
if desc.WireGuardPublicKey == "" || desc.WireGuardEndpoint == "" || desc.OverlayIPv4 == "" {
@@ -177,14 +211,14 @@ func peersForSnapshot(selfIdentityKey string, snapshot map[string]types.RelaySta
211
212
allowedIPs := []string{desc.OverlayIPv4 + "/32"}
213
allowedIPs = append(allowedIPs, desc.OverlayCIDRs...)
180
- peers = append(peers, types.DesiredPeer{
181
- WireGuardPublicKey: desc.WireGuardPublicKey,
182
- WireGuardEndpoint: desc.WireGuardEndpoint,
183
- AllowedIPs: allowedIPs,
214
+ peers = append(peers, desiredPeer{
215
+ wireGuardPublicKey: desc.WireGuardPublicKey,
216
+ wireGuardEndpoint: desc.WireGuardEndpoint,
217
+ allowedIPs: allowedIPs,
218
})
219
}
220
sort.Slice(peers, func(i, j int) bool {
187
- return peers[i].WireGuardPublicKey < peers[j].WireGuardPublicKey
221
+ return peers[i].wireGuardPublicKey < peers[j].wireGuardPublicKey
222
})
223
return peers
224
}
portal/wireguard/stack.go
+5
-6
@@ -15,7 +15,6 @@ import (
15
"golang.zx2c4.com/wireguard/device"
16
"golang.zx2c4.com/wireguard/tun/netstack"
17
18
- "github.com/gosuda/portal-tunnel/v2/types"
18
"github.com/gosuda/portal-tunnel/v2/utils"
19
)
20
@@ -121,7 +120,7 @@ func (s *stack) DialContext(ctx context.Context, network, address string) (net.C
120
return s.net.DialContextTCPAddrPort(ctx, netip.AddrPortFrom(ip, uint16(port)))
121
}
122
124
-func (s *stack) ApplyPeers(peers []types.DesiredPeer) error {
123
+func (s *stack) ApplyPeers(peers []desiredPeer) error {
124
if s == nil || s.device == nil {
125
return errors.New("wireguard is not initialized")
126
}
@@ -132,14 +131,14 @@ func (s *stack) ApplyPeers(peers []types.DesiredPeer) error {
131
nextPeerEndpoints := map[string]string{}
132
133
for _, peer := range peers {
135
- peerKey := strings.TrimSpace(peer.WireGuardPublicKey)
136
- publicKeyHex, err := utils.WireGuardKeyHex(peer.WireGuardPublicKey)
134
+ peerKey := strings.TrimSpace(peer.wireGuardPublicKey)
135
+ publicKeyHex, err := utils.WireGuardKeyHex(peer.wireGuardPublicKey)
136
if err != nil {
137
return fmt.Errorf("normalize peer %q public key: %w", peerKey, err)
138
}
139
140
resolvedEndpoint := ""
142
- if endpoint := peer.WireGuardEndpoint; endpoint != "" {
141
+ if endpoint := peer.wireGuardEndpoint; endpoint != "" {
142
resolvedEndpoint, err = resolvePeerEndpoint(endpoint)
143
if err != nil {
144
s.mu.Lock()
@@ -165,7 +164,7 @@ func (s *stack) ApplyPeers(peers []types.DesiredPeer) error {
164
nextPeerEndpoints[publicKeyHex] = resolvedEndpoint
165
}
166
168
- allowedIPs := utils.NormalizeIPPrefixes(peer.AllowedIPs)
167
+ allowedIPs := utils.NormalizeIPPrefixes(peer.allowedIPs)
168
for _, allowedIP := range allowedIPs {
169
builder.WriteString("allowed_ip=")
170
builder.WriteString(allowedIP)
sdk/expose.go
+127
-112
@@ -113,7 +113,10 @@ func Expose(ctx context.Context, cfg ExposeConfig) (*Exposure, error) {
113
}
114
115
if len(relayURLs) > 0 {
116
- exposure.relaySet.SetBootstrapRelayURLs(relayURLs)
116
+ if err := exposure.relaySet.SetBootstrapRelayURLs(relayURLs); err != nil {
117
+ _ = exposure.Close()
118
+ return nil, err
119
+ }
120
if err := exposure.reconcileRelayListeners(true); err != nil {
121
_ = exposure.Close()
122
return nil, err
@@ -121,11 +124,7 @@ func Expose(ctx context.Context, cfg ExposeConfig) (*Exposure, error) {
124
}
125
126
if cfg.Discovery {
124
- go func() {
125
- _ = exposure.relaySet.RunLoop(exposureCtx, exposure.rootCAPEM, func() error {
126
- return exposure.reconcileRelayListeners(false)
127
- })
128
- }()
127
+ go exposure.runDiscoveryLoop(exposureCtx)
128
}
129
go func() {
130
<-exposure.done
@@ -134,6 +133,31 @@ func Expose(ctx context.Context, cfg ExposeConfig) (*Exposure, error) {
133
134
return exposure, nil
135
}
136
+
137
+func (e *Exposure) runDiscoveryLoop(ctx context.Context) {
138
+ refresher, err := discovery.NewRefresher(e.relaySet, e.rootCAPEM, nil)
139
+ if err != nil {
140
+ return
141
+ }
142
+ ticker := time.NewTicker(types.DiscoveryPollInterval)
143
+ defer ticker.Stop()
144
+
145
+ for {
146
+ if err := refresher.Refresh(ctx); err != nil {
147
+ return
148
+ }
149
+ if err := e.reconcileRelayListeners(false); err != nil {
150
+ return
151
+ }
152
+
153
+ select {
154
+ case <-ctx.Done():
155
+ return
156
+ case <-ticker.C:
157
+ }
158
+ }
159
+}
160
+
161
func (e *Exposure) ActiveRelayURLs() []string {
162
return e.relaySet.ActiveRelayURLs()
163
}
@@ -149,6 +173,103 @@ func (e *Exposure) Identity() types.Identity {
173
return e.identity.Copy()
174
}
175
176
+func (e *Exposure) AcceptDatagram() (types.DatagramFrame, error) {
177
+ if !e.udpEnabled {
178
+ return types.DatagramFrame{}, net.ErrClosed
179
+ }
180
+
181
+ select {
182
+ case <-e.done:
183
+ return types.DatagramFrame{}, net.ErrClosed
184
+ case frame := <-e.datagrams:
185
+ return frame, nil
186
+ }
187
+}
188
+
189
+func (e *Exposure) SendDatagram(frame types.DatagramFrame) error {
190
+ if !e.udpEnabled {
191
+ return net.ErrClosed
192
+ }
193
+
194
+ e.listenerMu.RLock()
195
+ listener := e.relayListeners[frame.RelayURL]
196
+ e.listenerMu.RUnlock()
197
+ if listener == nil {
198
+ return net.ErrClosed
199
+ }
200
+ return listener.SendDatagram(frame)
201
+}
202
+
203
+func (e *Exposure) WaitDatagramReady(ctx context.Context) ([]string, error) {
204
+ if !e.udpEnabled {
205
+ return nil, errors.New("exposure does not have udp enabled")
206
+ }
207
+
208
+ ticker := time.NewTicker(50 * time.Millisecond)
209
+ defer ticker.Stop()
210
+
211
+ for {
212
+ e.listenerMu.RLock()
213
+ listeners := make([]*Listener, 0, len(e.relayListeners))
214
+ for _, listener := range e.relayListeners {
215
+ listeners = append(listeners, listener)
216
+ }
217
+ e.listenerMu.RUnlock()
218
+
219
+ addrs := make([]string, 0, len(listeners))
220
+ seen := make(map[string]struct{})
221
+ resolvedWithoutDatagram := true
222
+ for _, listener := range listeners {
223
+ if listener == nil {
224
+ continue
225
+ }
226
+
227
+ udpAddr, ready, pending := listener.DatagramReady()
228
+ if ready {
229
+ if _, ok := seen[udpAddr]; !ok {
230
+ seen[udpAddr] = struct{}{}
231
+ addrs = append(addrs, udpAddr)
232
+ }
233
+ }
234
+ if pending {
235
+ resolvedWithoutDatagram = false
236
+ }
237
+ }
238
+ if len(addrs) > 0 {
239
+ return addrs, nil
240
+ }
241
+ if resolvedWithoutDatagram {
242
+ return nil, errors.New("relay did not expose udp")
243
+ }
244
+
245
+ select {
246
+ case <-e.done:
247
+ return nil, net.ErrClosed
248
+ case <-ctx.Done():
249
+ return nil, ctx.Err()
250
+ case <-ticker.C:
251
+ }
252
+ }
253
+}
254
+
255
+func (e *Exposure) RunHTTP(ctx context.Context, handler http.Handler, localAddr string) error {
256
+ var relayListener net.Listener
257
+ e.listenerMu.RLock()
258
+ activeListeners := make([]*Listener, 0, len(e.relayListeners))
259
+ for _, relayURL := range e.relaySet.ActiveRelayURLs() {
260
+ listener, ok := e.relayListeners[relayURL]
261
+ if !ok {
262
+ continue
263
+ }
264
+ activeListeners = append(activeListeners, listener)
265
+ }
266
+ e.listenerMu.RUnlock()
267
+ if len(activeListeners) > 0 {
268
+ relayListener = e
269
+ }
270
+ return RunHTTP(ctx, relayListener, handler, localAddr)
271
+}
272
+
273
type exposureConn struct {
274
net.Conn
275
id uint64
@@ -243,13 +364,7 @@ func (e *Exposure) Close() error {
364
}
365
366
func (e *Exposure) reconcileRelayListeners(failOnError bool) error {
246
- if e.relaySet == nil {
247
- e.relaySet = discovery.NewRelaySet()
248
- }
367
e.listenerMu.Lock()
250
- if e.relayListeners == nil {
251
- e.relayListeners = make(map[string]*Listener)
252
- }
368
activeRelayURLs := e.relaySet.ActiveRelayURLs()
369
currentRelayURLs := make([]string, 0, len(e.relayListeners))
370
for relayURL := range e.relayListeners {
@@ -298,9 +413,6 @@ func (e *Exposure) reconcileRelayListeners(failOnError bool) error {
413
}
414
415
e.listenerMu.Lock()
301
- if e.relayListeners == nil {
302
- e.relayListeners = make(map[string]*Listener, 1)
303
- }
416
if _, exists := e.relayListeners[relayURL]; exists {
417
e.listenerMu.Unlock()
418
_ = listener.Close()
@@ -375,100 +487,3 @@ func (e *Exposure) runListenerAcceptLoop(listener *Listener) {
487
}
488
}
489
}
378
-
379
-func (e *Exposure) AcceptDatagram() (types.DatagramFrame, error) {
380
- if !e.udpEnabled {
381
- return types.DatagramFrame{}, net.ErrClosed
382
- }
383
-
384
- select {
385
- case <-e.done:
386
- return types.DatagramFrame{}, net.ErrClosed
387
- case frame := <-e.datagrams:
388
- return frame, nil
389
- }
390
-}
391
-
392
-func (e *Exposure) SendDatagram(frame types.DatagramFrame) error {
393
- if !e.udpEnabled {
394
- return net.ErrClosed
395
- }
396
-
397
- e.listenerMu.RLock()
398
- listener := e.relayListeners[frame.RelayURL]
399
- e.listenerMu.RUnlock()
400
- if listener == nil {
401
- return net.ErrClosed
402
- }
403
- return listener.SendDatagram(frame)
404
-}
405
-
406
-func (e *Exposure) WaitDatagramReady(ctx context.Context) ([]string, error) {
407
- if !e.udpEnabled {
408
- return nil, errors.New("exposure does not have udp enabled")
409
- }
410
-
411
- ticker := time.NewTicker(50 * time.Millisecond)
412
- defer ticker.Stop()
413
-
414
- for {
415
- e.listenerMu.RLock()
416
- listeners := make([]*Listener, 0, len(e.relayListeners))
417
- for _, listener := range e.relayListeners {
418
- listeners = append(listeners, listener)
419
- }
420
- e.listenerMu.RUnlock()
421
-
422
- addrs := make([]string, 0, len(listeners))
423
- seen := make(map[string]struct{})
424
- resolvedWithoutDatagram := true
425
- for _, listener := range listeners {
426
- if listener == nil {
427
- continue
428
- }
429
-
430
- udpAddr, ready, pending := listener.DatagramReady()
431
- if ready {
432
- if _, ok := seen[udpAddr]; !ok {
433
- seen[udpAddr] = struct{}{}
434
- addrs = append(addrs, udpAddr)
435
- }
436
- }
437
- if pending {
438
- resolvedWithoutDatagram = false
439
- }
440
- }
441
- if len(addrs) > 0 {
442
- return addrs, nil
443
- }
444
- if resolvedWithoutDatagram {
445
- return nil, errors.New("relay did not expose udp")
446
- }
447
-
448
- select {
449
- case <-e.done:
450
- return nil, net.ErrClosed
451
- case <-ctx.Done():
452
- return nil, ctx.Err()
453
- case <-ticker.C:
454
- }
455
- }
456
-}
457
-
458
-func (e *Exposure) RunHTTP(ctx context.Context, handler http.Handler, localAddr string) error {
459
- var relayListener net.Listener
460
- e.listenerMu.RLock()
461
- activeListeners := make([]*Listener, 0, len(e.relayListeners))
462
- for _, relayURL := range e.relaySet.ActiveRelayURLs() {
463
- listener, ok := e.relayListeners[relayURL]
464
- if !ok {
465
- continue
466
- }
467
- activeListeners = append(activeListeners, listener)
468
- }
469
- e.listenerMu.RUnlock()
470
- if len(activeListeners) > 0 {
471
- relayListener = e
472
- }
473
- return RunHTTP(ctx, relayListener, handler, localAddr)
474
-}
sdk/expose_test.go
+25
-8
@@ -30,6 +30,15 @@ func mustRelayDescriptor(t *testing.T, relayName, relayURL string) types.RelayDe
30
return desc
31
}
32
33
+func applyRelayDiscovery(t *testing.T, set *discovery.RelaySet, identity types.Identity, targetURL string, resp types.DiscoveryResponse, now time.Time) error {
34
+ t.Helper()
35
+ _, warnErr, err := set.ApplyRelayDiscoveryResponse(identity, targetURL, resp, now)
36
+ if err != nil {
37
+ return err
38
+ }
39
+ return warnErr
40
+}
41
+
42
func TestExposureBanRelayURLMovesRelay(t *testing.T) {
43
const (
44
relayA = "https://relay-a.example"
@@ -49,13 +58,15 @@ func TestExposureBanRelayURLMovesRelay(t *testing.T) {
58
relaySet: discovery.NewRelaySet(),
59
relayListeners: make(map[string]*Listener, 2),
60
}
52
- exposure.relaySet.SetBootstrapRelayURLs([]string{relayA, relayB})
61
+ if err := exposure.relaySet.SetBootstrapRelayURLs([]string{relayA, relayB}); err != nil {
62
+ t.Fatalf("SetBootstrapRelayURLs() error = %v", err)
63
+ }
64
exposure.relayListeners = map[string]*Listener{
65
relayA: listener,
66
relayB: {},
67
}
68
58
- exposure.relaySet.BanRelayURL(relayA, "test")
69
+ exposure.relaySet.BanRelayURL(relayA)
70
exposure.listenerMu.Lock()
71
delete(exposure.relayListeners, relayA)
72
exposure.listenerMu.Unlock()
@@ -86,12 +97,14 @@ func TestExposureSetRelayURLsSkipsBannedRelay(t *testing.T) {
97
relaySet: discovery.NewRelaySet(),
98
relayListeners: make(map[string]*Listener, 1),
99
}
89
- exposure.relaySet.BanRelayURL(relayB, "test")
100
+ exposure.relaySet.BanRelayURL(relayB)
101
exposure.relayListeners = map[string]*Listener{
102
relayA: {},
103
}
104
94
- exposure.relaySet.SetBootstrapRelayURLs([]string{relayA, relayB})
105
+ if err := exposure.relaySet.SetBootstrapRelayURLs([]string{relayA, relayB}); err != nil {
106
+ t.Fatalf("SetBootstrapRelayURLs() error = %v", err)
107
+ }
108
if err := exposure.reconcileRelayListeners(false); err != nil {
109
t.Fatalf("reconcileRelayListeners() error = %v", err)
110
}
@@ -124,7 +137,9 @@ func TestExposureSetRelayURLsRemovesStaleListener(t *testing.T) {
137
relaySet: discovery.NewRelaySet(),
138
relayListeners: make(map[string]*Listener, 2),
139
}
127
- exposure.relaySet.SetBootstrapRelayURLs([]string{relayA, relayB})
140
+ if err := exposure.relaySet.SetBootstrapRelayURLs([]string{relayA, relayB}); err != nil {
141
+ t.Fatalf("SetBootstrapRelayURLs() error = %v", err)
142
+ }
143
exposure.relayListeners = map[string]*Listener{
144
relayA: {
145
api: &apiClient{baseURL: relayAURL},
@@ -136,7 +151,9 @@ func TestExposureSetRelayURLsRemovesStaleListener(t *testing.T) {
151
},
152
}
153
139
- exposure.relaySet.SetBootstrapRelayURLs([]string{relayB})
154
+ if err := exposure.relaySet.SetBootstrapRelayURLs([]string{relayB}); err != nil {
155
+ t.Fatalf("SetBootstrapRelayURLs() error = %v", err)
156
+ }
157
if err := exposure.reconcileRelayListeners(false); err != nil {
158
t.Fatalf("reconcileRelayListeners() error = %v", err)
159
}
@@ -167,12 +184,12 @@ func TestExposurePinDiscoveredDescriptorAllowsURLChangeForSameIdentity(t *testin
184
exposure := &Exposure{relaySet: discovery.NewRelaySet()}
185
desc := mustRelayDescriptor(t, "relay-a", "https://relay-a.example")
186
170
- if err := exposure.relaySet.ApplyRelayDiscoveryResponseSimple(desc.Identity, desc.APIHTTPSAddr, types.DiscoveryResponse{ProtocolVersion: types.ProtocolVersion, Self: desc}, time.Now().UTC()); err != nil {
187
+ if err := applyRelayDiscovery(t, exposure.relaySet, desc.Identity, desc.APIHTTPSAddr, types.DiscoveryResponse{ProtocolVersion: types.ProtocolVersion, Self: desc}, time.Now().UTC()); err != nil {
188
t.Fatalf("ApplyRelayDiscoveryResponse() error = %v", err)
189
}
190
191
changedURL := mustRelayDescriptor(t, desc.Name, "https://relay-b.example")
175
- err := exposure.relaySet.ApplyRelayDiscoveryResponseSimple(desc.Identity, "", types.DiscoveryResponse{ProtocolVersion: types.ProtocolVersion, Self: changedURL}, time.Now().UTC())
192
+ err := applyRelayDiscovery(t, exposure.relaySet, desc.Identity, "", types.DiscoveryResponse{ProtocolVersion: types.ProtocolVersion, Self: changedURL}, time.Now().UTC())
193
if err != nil {
194
t.Fatalf("ApplyRelayDiscoveryResponse() error = %v, want nil for same relay identity", err)
195
}
sdk/listener.go
+1
-1
@@ -546,7 +546,7 @@ func (l *Listener) closed() bool {
546
547
func (l *Listener) ban() {
548
if l.relaySet != nil && l.api != nil && l.api.baseURL != nil {
549
- l.relaySet.BanRelayURL(l.api.baseURL.String(), "manual")
549
+ l.relaySet.BanRelayURL(l.api.baseURL.String())
550
}
551
_ = l.Close()
552
}
sdk/mitm_test.go
+6
-2
@@ -221,7 +221,9 @@ func TestMITMProbeDetectionBansListener(t *testing.T) {
221
registered: make(chan struct{}),
222
banMITM: true,
223
}
224
- listener.relaySet.SetBootstrapRelayURLs([]string{relayURL.String()})
224
+ if err := listener.relaySet.SetBootstrapRelayURLs([]string{relayURL.String()}); err != nil {
225
+ t.Fatalf("SetBootstrapRelayURLs() error = %v", err)
226
+ }
227
listener.mitmManager = newMITMManager(context.Background(), listener)
228
229
listener.mitmManager.logResult(MITMProbeReport{
@@ -254,7 +256,9 @@ func TestMITMProbeDetectionWarnsWithoutBanningListener(t *testing.T) {
256
registered: make(chan struct{}),
257
banMITM: false,
258
}
257
- listener.relaySet.SetBootstrapRelayURLs([]string{relayURL.String()})
259
+ if err := listener.relaySet.SetBootstrapRelayURLs([]string{relayURL.String()}); err != nil {
260
+ t.Fatalf("SetBootstrapRelayURLs() error = %v", err)
261
+ }
262
listener.mitmManager = newMITMManager(context.Background(), listener)
263
264
listener.mitmManager.logResult(MITMProbeReport{
types/overlay.go
deleted
-10
@@ -1,10 +0,0 @@
1
-package types
2
-
3
-// DesiredPeer describes a WireGuard peer that should be programmed into the
4
-// local runtime.
5
-type DesiredPeer struct {
6
- RelayID string `json:"relay_id"`
7
- WireGuardPublicKey string `json:"wireguard_public_key"`
8
- WireGuardEndpoint string `json:"wireguard_endpoint"`
9
- AllowedIPs []string `json:"allowed_ips,omitempty"`
10
-}
types/relay_state.go
deleted
-16
@@ -1,16 +0,0 @@
1
-package types
2
-
3
-import "time"
4
-
5
-// RelayState captures the last-known descriptor and local state for a relay
6
-// observed through discovery.
7
-type RelayState struct {
8
- Descriptor RelayDescriptor `json:"descriptor"`
9
- Bootstrap bool `json:"bootstrap,omitempty"`
10
- Advertised bool `json:"advertised,omitempty"`
11
- Expired bool `json:"expired,omitempty"`
12
- Banned bool `json:"banned,omitempty"`
13
- FirstSeenAt time.Time `json:"first_seen_at"`
14
- LastSeenAt time.Time `json:"last_seen_at"`
15
- ConsecutiveFailures int `json:"consecutive_failures,omitempty"`
16
-}