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 -}