refact: WIP separate policy from relayset

Chang committed Apr 12, 2026 at 16:32 UTC b211765c33af5bbeff035e18cc1aabed14e87d9f
9 files changed +513 -673
portal/discovery/policy.go
+267 -10
@@ -2,27 +2,284 @@ package discovery
2
3 import (
4 "errors"
5 + "fmt"
6 + "net/http"
7 + "sort"
8 + "strings"
9 "time"
10
11 "github.com/gosuda/portal-tunnel/v2/types"
12 + "github.com/gosuda/portal-tunnel/v2/utils"
13 )
14
15 +type RelayStatus uint8
16 +
17 const (
11 - DiscoveryPollInterval = 1 * time.Minute
18 + RelayStatusUnknown RelayStatus = iota
19 + RelayStatusHinted
20 + RelayStatusConfirmed
21 + RelayStatusInvalid
22 + RelayStatusExpired
23 + RelayStatusBanned
24 )
25
14 -func RequireOverlayRelayDescriptor(desc types.RelayDescriptor) error {
15 - if !desc.SupportsOverlayPeer {
16 - return errors.New("descriptor does not support overlay peer")
26 +type RelayState struct {
27 + Descriptor types.RelayDescriptor
28 + Bootstrap bool
29 + FirstSeenAt time.Time
30 + LastSeenAt time.Time
31 + Status RelayStatus
32 + Active bool
33 + Advertise bool
34 + BootstrapDiscovery bool
35 + DirectDiscovery bool
36 + OverlayDiscovery bool
37 + OverlayPeer bool
38 +
39 + consecutiveFailures int
40 +}
41 +
42 +func newRelayState(desc types.RelayDescriptor, seenAt time.Time) (RelayState, error) {
43 + state := RelayState{
44 + Descriptor: desc,
45 + }
46 + if seenAt.IsZero() {
47 + return state, nil
48 + }
49 +
50 + seenAt = seenAt.UTC()
51 + normalized, err := utils.NormalizeDescriptor(desc)
52 + if err != nil {
53 + return RelayState{}, err
54 }
18 - if desc.WireGuardPublicKey == "" {
19 - return errors.New("descriptor wireguard public key is required")
55 +
56 + switch {
57 + case normalized.Name == "":
58 + return RelayState{}, errors.New("identity.name is required")
59 + case normalized.APIHTTPSAddr == "":
60 + return RelayState{}, errors.New("api_https_addr is required")
61 + case normalized.RelayID == "":
62 + return RelayState{}, errors.New("relay_id is required")
63 + case normalized.RelayID != normalized.APIHTTPSAddr:
64 + return RelayState{}, errors.New("relay_id must match api_https_addr")
65 + case normalized.Sequence == 0:
66 + return RelayState{}, errors.New("sequence is required")
67 + case normalized.Version == 0:
68 + return RelayState{}, errors.New("version is required")
69 + case normalized.IssuedAt.IsZero():
70 + return RelayState{}, errors.New("issued_at is required")
71 + case normalized.ExpiresAt.IsZero():
72 + return RelayState{}, errors.New("expires_at is required")
73 + case normalized.ExpiresAt.Before(seenAt):
74 + return RelayState{}, errors.New("descriptor expired")
75 + case normalized.IssuedAt.After(normalized.ExpiresAt):
76 + return RelayState{}, errors.New("issued_at must be before expires_at")
77 }
21 - if desc.WireGuardEndpoint == "" {
22 - return errors.New("descriptor wireguard endpoint is required")
78 +
79 + state.Descriptor = normalized
80 + state.FirstSeenAt = seenAt
81 + state.LastSeenAt = seenAt
82 + return state, nil
83 +}
84 +
85 +func newRelayHintState(relayURL string) RelayState {
86 + state, _ := newRelayState(
87 + types.RelayDescriptor{
88 + Identity: types.Identity{
89 + Name: utils.PortalRootHost(relayURL),
90 + },
91 + RelayID: relayURL,
92 + APIHTTPSAddr: relayURL,
93 + Version: 1,
94 + },
95 + time.Time{},
96 + )
97 + return state
98 +}
99 +
100 +func (DefaultRelayPolicy) DiscoveryStates(targetIdentity types.Identity, targetURL string, selfRelayKey string, selfRelayURL string, resp types.DiscoveryResponse, seenAt time.Time) (RelayState, []RelayState, error) {
101 + protocolVersion := strings.TrimSpace(resp.ProtocolVersion)
102 + if protocolVersion != types.ProtocolVersion {
103 + return RelayState{}, nil, fmt.Errorf("relay protocol version mismatch: relay=%q client=%q", protocolVersion, types.ProtocolVersion)
104 }
24 - if desc.OverlayIPv4 == "" {
25 - return errors.New("descriptor overlay ipv4 is required")
105 +
106 + selfState, err := newRelayState(resp.Self, seenAt)
107 + if err != nil {
108 + return RelayState{}, nil, err
109 + }
110 + if err := requireTargetRelayDescriptor(selfState.Descriptor, targetIdentity, targetURL); err != nil {
111 + return RelayState{}, nil, err
112 + }
113 +
114 + relayStates := make([]RelayState, 0, len(resp.Relays))
115 + seen := map[string]struct{}{selfState.Descriptor.Key(): {}}
116 + for _, descriptor := range resp.Relays {
117 + relayState, err := newRelayState(descriptor, seenAt)
118 + if err != nil {
119 + continue
120 + }
121 + if selfRelay(relayState, selfRelayKey, selfRelayURL) {
122 + continue
123 + }
124 + relayKey := relayState.Descriptor.Key()
125 + if _, ok := seen[relayKey]; ok {
126 + continue
127 + }
128 + seen[relayKey] = struct{}{}
129 + relayStates = append(relayStates, relayState)
130 + }
131 + return selfState, relayStates, nil
132 +}
133 +
134 +func selfRelay(state RelayState, selfRelayKey string, selfRelayURL string) bool {
135 + desc := state.Descriptor
136 + relayKey := desc.Key()
137 + return (relayKey != "" && selfRelayKey != "" && relayKey == selfRelayKey) ||
138 + (selfRelayURL != "" && desc.APIHTTPSAddr == selfRelayURL)
139 +}
140 +
141 +func requireTargetRelayDescriptor(desc types.RelayDescriptor, targetIdentity types.Identity, targetURL string) error {
142 + if strings.TrimSpace(targetIdentity.Name) == "" && strings.TrimSpace(targetIdentity.Address) == "" {
143 + return errors.New("target relay identity is required")
144 + }
145 + targetName := strings.TrimSpace(targetIdentity.Name)
146 + if targetName != "" {
147 + normalizedTargetName := utils.NormalizeHostname(targetName)
148 + if desc.Name != normalizedTargetName {
149 + return errors.New("descriptor name does not match target relay")
150 + }
151 + }
152 + targetAddress := strings.TrimSpace(targetIdentity.Address)
153 + if targetAddress != "" {
154 + normalizedTargetAddress, err := utils.NormalizeEVMAddress(targetAddress)
155 + if err != nil {
156 + return err
157 + }
158 + if desc.Address != normalizedTargetAddress {
159 + return errors.New("descriptor address does not match target relay")
160 + }
161 + }
162 + if targetURL != "" {
163 + normalizedTargetURL, err := utils.NormalizeRelayURL(targetURL)
164 + if err != nil {
165 + return err
166 + }
167 + if desc.APIHTTPSAddr != normalizedTargetURL {
168 + return errors.New("descriptor api_https_addr does not match target url")
169 + }
170 }
171 return nil
172 }
173 +
174 +func (state RelayState) hasDescriptor() bool {
175 + return !state.LastSeenAt.IsZero() && state.Descriptor.Key() != "" && state.Descriptor.APIHTTPSAddr != ""
176 +}
177 +
178 +type RelayPolicy interface {
179 + DiscoveryStates(types.Identity, string, string, string, types.DiscoveryResponse, time.Time) (RelayState, []RelayState, error)
180 + Decide(RelayState) RelayState
181 + OnDiscovered(RelayState, bool) RelayState
182 + OnFailure(RelayState, error, int) (RelayState, bool, string)
183 + OnBanned(RelayState) RelayState
184 + KeepState(RelayState) bool
185 + AdvertisedDescriptors([]RelayState) []types.RelayDescriptor
186 +}
187 +
188 +type DefaultRelayPolicy struct{}
189 +
190 +func (DefaultRelayPolicy) Decide(state RelayState) RelayState {
191 + now := time.Now().UTC()
192 + status := state.Status
193 + switch {
194 + case status == RelayStatusBanned || status == RelayStatusExpired:
195 + case !state.Descriptor.ExpiresAt.IsZero() && !state.Descriptor.ExpiresAt.After(now):
196 + status = RelayStatusExpired
197 + case state.Descriptor.SupportsOverlayPeer &&
198 + (state.Descriptor.WireGuardPublicKey == "" ||
199 + state.Descriptor.WireGuardEndpoint == "" ||
200 + state.Descriptor.OverlayIPv4 == ""):
201 + status = RelayStatusInvalid
202 + case status == RelayStatusUnknown:
203 + status = RelayStatusHinted
204 + }
205 +
206 + usable := status == RelayStatusHinted || status == RelayStatusConfirmed
207 + state.Status = status
208 + state.Active = (state.Bootstrap && usable) || status == RelayStatusConfirmed
209 + state.Advertise = status == RelayStatusConfirmed
210 + state.BootstrapDiscovery = state.Bootstrap && usable
211 + state.DirectDiscovery = !state.Bootstrap && usable
212 + state.OverlayDiscovery = !state.Bootstrap && usable && state.Descriptor.SupportsOverlayPeer
213 + state.OverlayPeer = usable && state.Descriptor.SupportsOverlayPeer
214 + return state
215 +}
216 +
217 +func (DefaultRelayPolicy) OnDiscovered(state RelayState, advertise bool) RelayState {
218 + if state.Status == RelayStatusBanned {
219 + return state
220 + }
221 + if advertise {
222 + state.Status = RelayStatusConfirmed
223 + state.consecutiveFailures = 0
224 + return state
225 + }
226 + if state.Status != RelayStatusConfirmed {
227 + state.Status = RelayStatusHinted
228 + state.consecutiveFailures = 0
229 + }
230 + return state
231 +}
232 +
233 +func (DefaultRelayPolicy) OnFailure(state RelayState, err error, recoveryFailures int) (RelayState, bool, string) {
234 + if state.Status == RelayStatusBanned {
235 + return state, false, ""
236 + }
237 + state.consecutiveFailures++
238 + if state.Status != RelayStatusExpired && state.consecutiveFailures >= recoveryFailures {
239 + state.Status = RelayStatusExpired
240 + return state, true, "recovery"
241 + }
242 + var apiErr *types.APIRequestError
243 + if errors.As(err, &apiErr) &&
244 + (apiErr.StatusCode == http.StatusForbidden ||
245 + apiErr.StatusCode == http.StatusNotFound ||
246 + apiErr.StatusCode == http.StatusGone) {
247 + state.Status = RelayStatusExpired
248 + return state, true, "status"
249 + }
250 + return state, false, ""
251 +}
252 +
253 +func (DefaultRelayPolicy) OnBanned(state RelayState) RelayState {
254 + state.Status = RelayStatusBanned
255 + return state
256 +}
257 +
258 +func (DefaultRelayPolicy) KeepState(state RelayState) bool {
259 + return state.hasDescriptor() ||
260 + state.Status == RelayStatusBanned ||
261 + state.Status == RelayStatusExpired ||
262 + state.consecutiveFailures > 0
263 +}
264 +
265 +func (p DefaultRelayPolicy) AdvertisedDescriptors(states []RelayState) []types.RelayDescriptor {
266 + if len(states) == 0 {
267 + return nil
268 + }
269 +
270 + out := make([]types.RelayDescriptor, 0, len(states))
271 + for _, state := range states {
272 + state = p.Decide(state)
273 + if state.Descriptor.APIHTTPSAddr == "" || !state.Advertise {
274 + continue
275 + }
276 + out = append(out, state.Descriptor)
277 + }
278 + if len(out) == 0 {
279 + return nil
280 + }
281 + sort.Slice(out, func(i, j int) bool {
282 + return out[i].APIHTTPSAddr < out[j].APIHTTPSAddr
283 + })
284 + return out
285 +}
portal/discovery/refresher.go
+42 -36
@@ -16,12 +16,13 @@ import (
16 )
17
18 const (
19 + DiscoveryPollInterval = 1 * time.Minute
20 defaultRecoveryFailures = 3
21 )
22
23 type OverlayRuntime interface {
24 DiscoverRelay(context.Context, types.RelayDescriptor) (types.DiscoveryResponse, error)
24 - Sync(map[string]RelayState) error
25 + Sync([]RelayState) error
26 }
27
28 type Refresher struct {
@@ -67,7 +68,7 @@ func (r *Refresher) Refresh(ctx context.Context) error {
68 if r.overlay == nil {
69 return ctx.Err()
70 }
70 - if err := r.overlay.Sync(r.relaySet.View()); err != nil {
71 + if err := r.overlay.Sync(r.relaySet.RelayStates()); err != nil {
72 log.Warn().
73 Err(err).
74 Msg("sync wireguard peers")
@@ -77,8 +78,13 @@ func (r *Refresher) Refresh(ctx context.Context) error {
78 }
79
80 func (r *Refresher) refreshHTTPS(ctx context.Context) error {
80 - for _, bootstrap := range r.relaySet.BootstrapDescriptors() {
81 - resp, err := r.discoverHTTPS(ctx, bootstrap)
81 + states := r.relaySet.RelayStates()
82 + for _, state := range states {
83 + if state.Descriptor.APIHTTPSAddr == "" || !state.BootstrapDiscovery {
84 + continue
85 + }
86 + relay := state.Descriptor
87 + resp, err := r.discoverHTTPS(ctx, relay)
88 if err != nil {
89 if ctx.Err() != nil {
90 return ctx.Err()
@@ -87,7 +93,7 @@ func (r *Refresher) refreshHTTPS(ctx context.Context) error {
93 }
94
95 now := time.Now().UTC()
90 - _, _, err = r.relaySet.ApplyRelayDiscoveryResponse(bootstrap.Identity, bootstrap.APIHTTPSAddr, resp, now)
96 + _, err = r.relaySet.ApplyRelayDiscoveryResponse(relay.Identity, relay.APIHTTPSAddr, resp, now)
97 if err != nil {
98 continue
99 }
@@ -96,10 +102,15 @@ func (r *Refresher) refreshHTTPS(ctx context.Context) error {
102 return err
103 }
104
99 - for _, relay := range r.relaySet.confirmableDescriptors() {
100 - if r.overlay != nil && relay.SupportsOverlayPeer {
105 + states = r.relaySet.RelayStates()
106 + for _, state := range states {
107 + if state.Descriptor.APIHTTPSAddr == "" || !state.DirectDiscovery {
108 + continue
109 + }
110 + if r.overlay != nil && state.Descriptor.SupportsOverlayPeer {
111 continue
112 }
113 + relay := state.Descriptor
114 resp, err := r.discoverHTTPS(ctx, relay)
115 if err != nil {
116 if ctx.Err() != nil {
@@ -110,7 +121,7 @@ func (r *Refresher) refreshHTTPS(ctx context.Context) error {
121 }
122
123 now := time.Now().UTC()
113 - _, _, err = r.relaySet.ApplyRelayDiscoveryResponse(relay.Identity, relay.APIHTTPSAddr, resp, now)
124 + _, err = r.relaySet.ApplyRelayDiscoveryResponse(relay.Identity, relay.APIHTTPSAddr, resp, now)
125 if err != nil {
126 r.logDirectDiscoveryFailure(relay, err, r.directRecoveryFailures)
127 continue
@@ -133,46 +144,41 @@ func (r *Refresher) discoverHTTPS(ctx context.Context, relay types.RelayDescript
144 }
145
146 func (r *Refresher) refreshOverlay(ctx context.Context) error {
136 - for _, relay := range r.relaySet.SyncableDescriptors() {
147 + for _, state := range r.relaySet.RelayStates() {
148 + if state.Descriptor.APIHTTPSAddr == "" || !state.OverlayDiscovery {
149 + continue
150 + }
151 +
152 + relay := state.Descriptor
153 var failureErr error
154
139 - if err := RequireOverlayRelayDescriptor(relay); err != nil {
155 + resp, err := r.overlay.DiscoverRelay(ctx, relay)
156 + if err != nil {
157 + if ctx.Err() != nil {
158 + return ctx.Err()
159 + }
160 failureErr = err
161 } else {
142 - resp, err := r.overlay.DiscoverRelay(ctx, relay)
143 - if err != nil {
144 - if ctx.Err() != nil {
145 - return ctx.Err()
162 + now := time.Now().UTC()
163 + relaySetChanged, err := r.relaySet.ApplyRelayDiscoveryResponse(relay.Identity, relay.APIHTTPSAddr, resp, now)
164 + if relaySetChanged {
165 + if syncErr := r.overlay.Sync(r.relaySet.RelayStates()); syncErr != nil {
166 + log.Warn().
167 + Err(syncErr).
168 + Str("relay", relay.APIHTTPSAddr).
169 + Msg("sync wireguard peers")
170 }
171 + }
172 + if err != nil {
173 failureErr = err
174 } else {
149 - now := time.Now().UTC()
150 - relaySetChanged, warnErr, err := r.relaySet.ApplyOverlayRelayDiscoveryResponse(relay.Identity, relay.APIHTTPSAddr, resp, now)
151 - if relaySetChanged {
152 - view := r.relaySet.View()
153 - if syncErr := r.overlay.Sync(view); syncErr != nil {
154 - if warnErr == nil {
155 - warnErr = syncErr
156 - }
157 - }
158 - }
159 - if err != nil {
160 - failureErr = err
161 - } else {
162 - if warnErr != nil {
163 - log.Warn().
164 - Err(warnErr).
165 - Str("relay", relay.APIHTTPSAddr).
166 - Msg("overlay relay discovery completed with warnings")
167 - }
168 - continue
169 - }
175 + continue
176 }
177 }
178
179 expired, expireReason, consecutiveFailures := r.relaySet.RecordDiscoveryFailure(relay.Identity, relay.APIHTTPSAddr, failureErr, r.overlayRecoveryFailures)
180 if expired {
175 - if syncErr := r.overlay.Sync(r.relaySet.View()); syncErr != nil && failureErr == nil {
181 + if syncErr := r.overlay.Sync(r.relaySet.RelayStates()); syncErr != nil && failureErr == nil {
182 failureErr = syncErr
183 }
184 }
portal/discovery/relayset.go
+151 -592
@@ -2,8 +2,6 @@ package discovery
2
3 import (
4 "errors"
5 - "fmt"
6 - "net/http"
5 "reflect"
6 "slices"
7 "sort"
@@ -11,76 +9,36 @@ import (
9 "sync"
10 "time"
11
14 - "github.com/rs/zerolog/log"
15 -
12 "github.com/gosuda/portal-tunnel/v2/types"
13 "github.com/gosuda/portal-tunnel/v2/utils"
14 )
15
20 -type relayStatus uint8
21 -
22 -const (
23 - relayStatusHinted relayStatus = iota
24 - relayStatusConfirmed
25 - relayStatusExpired
26 -)
27 -
28 -// RelayState is the single relay shape shared by discovery storage and overlay sync.
29 -// Package-private fields are internal-only local state.
30 -type RelayState struct {
31 - Descriptor types.RelayDescriptor
32 - FirstSeenAt time.Time
33 - LastSeenAt time.Time
34 - Banned bool
35 - Expired bool
36 - status relayStatus
37 - consecutiveFailures int
38 -}
39 -
40 -// RelayCandidate is a relay fact projected for client-side selection.
41 -// Discovery owns the facts; callers decide how many candidates to use.
42 -type RelayCandidate struct {
43 - Descriptor types.RelayDescriptor
44 - Bootstrap bool
45 - FirstSeenAt time.Time
46 - LastSeenAt time.Time
47 -}
48 -
49 -type SelectionOptions struct {
50 - Limit int
51 -}
52 -
53 -func (s RelayState) isDefaultLocalState() bool {
54 - return !s.Banned && s.status == relayStatusHinted && s.consecutiveFailures == 0
55 -}
56 -
16 // RelaySet owns the shared relay discovery view: configured bootstrap relay URLs,
17 // the latest validated descriptor seen for each relay, and local runtime state
18 // such as ban/reachability/failure tracking.
19 type RelaySet struct {
20 mu sync.RWMutex
21 knownRelayURLs []string
63 - relayKeysByURL map[string]string
22 relays map[string]RelayState
65 - localByURL map[string]RelayState
23 + policy RelayPolicy
24 selfRelayKey string
25 selfRelayURL string
26 }
27
70 -type relayDescriptorProjection struct {
71 - state RelayState
72 - relayURL string
73 - bootstrap bool
74 -}
75 -
28 func NewRelaySet(identity types.Identity, relayURL string, bootstrapRelayURLs []string) (*RelaySet, error) {
77 - set := &RelaySet{
78 - relayKeysByURL: make(map[string]string),
79 - relays: make(map[string]RelayState),
80 - localByURL: make(map[string]RelayState),
29 + if relayURL != "" {
30 + normalized, err := utils.NormalizeRelayURL(relayURL)
31 + if err != nil {
32 + return nil, err
33 + }
34 + relayURL = normalized
35 }
82 - if err := set.SetSelfRelay(identity, relayURL); err != nil {
83 - return nil, err
36 +
37 + set := &RelaySet{
38 + relays: make(map[string]RelayState),
39 + policy: DefaultRelayPolicy{},
40 + selfRelayKey: identity.Key(),
41 + selfRelayURL: relayURL,
42 }
43 if err := set.SetBootstrapRelayURLs(bootstrapRelayURLs); err != nil {
44 return nil, err
@@ -88,180 +46,81 @@ func NewRelaySet(identity types.Identity, relayURL string, bootstrapRelayURLs []
46 return set, nil
47 }
48
91 -func relayExpiredAt(state RelayState, now time.Time) bool {
92 - if state.status == relayStatusExpired {
93 - return true
94 - }
95 - if state.Descriptor.ExpiresAt.IsZero() {
96 - return false
97 - }
98 - if now.IsZero() {
99 - now = time.Now().UTC()
49 +func (s *RelaySet) SetRelayPolicy(policy RelayPolicy) {
50 + if policy == nil {
51 + policy = DefaultRelayPolicy{}
52 }
101 - return !state.Descriptor.ExpiresAt.After(now)
102 -}
103 -
104 -func (s *RelaySet) isSelfRelayURLLocked(relayURL string) bool {
105 - relayURL = strings.TrimSpace(relayURL)
106 - return relayURL != "" && s.selfRelayURL != "" && relayURL == s.selfRelayURL
107 -}
108 -
109 -func (s *RelaySet) isSelfRelayDescriptorLocked(desc types.RelayDescriptor) bool {
110 - if relayKey := desc.Key(); relayKey != "" && s.selfRelayKey != "" && relayKey == s.selfRelayKey {
111 - return true
112 - }
113 - return s.isSelfRelayURLLocked(desc.APIHTTPSAddr)
114 -}
115 -
116 -func (s *RelaySet) storeLocalStateLocked(relayURL string, state RelayState) {
117 - relayURL = strings.TrimSpace(relayURL)
118 - if relayURL == "" {
119 - return
120 - }
121 - if state.isDefaultLocalState() {
122 - delete(s.localByURL, relayURL)
123 - return
124 - }
125 - s.localByURL[relayURL] = state
53 + s.mu.Lock()
54 + defer s.mu.Unlock()
55 + s.policy = policy
56 }
57
128 -func (s *RelaySet) SetSelfRelay(identity types.Identity, relayURL string) error {
129 - relayURL = strings.TrimSpace(relayURL)
130 - if relayURL != "" {
131 - normalized, err := utils.NormalizeRelayURL(relayURL)
132 - if err != nil {
133 - return err
134 - }
135 - relayURL = normalized
58 +func (s *RelaySet) SetBootstrapRelayURLs(inputs []string) error {
59 + normalized, err := utils.NormalizeRelayURLs(inputs...)
60 + if err != nil {
61 + return err
62 }
63
64 s.mu.Lock()
65 defer s.mu.Unlock()
140 - s.selfRelayKey = identity.Key()
141 - s.selfRelayURL = relayURL
142 - if s.selfRelayURL != "" {
143 - filtered := s.knownRelayURLs[:0]
144 - for _, knownRelayURL := range s.knownRelayURLs {
145 - if s.isSelfRelayURLLocked(knownRelayURL) {
146 - continue
147 - }
148 - filtered = append(filtered, knownRelayURL)
149 - }
150 - s.knownRelayURLs = filtered
151 - delete(s.localByURL, s.selfRelayURL)
152 - delete(s.relayKeysByURL, s.selfRelayURL)
153 - }
154 - if s.selfRelayKey != "" {
155 - if record, ok := s.relays[s.selfRelayKey]; ok {
156 - delete(s.localByURL, record.Descriptor.APIHTTPSAddr)
157 - delete(s.relayKeysByURL, record.Descriptor.APIHTTPSAddr)
158 - }
159 - delete(s.relays, s.selfRelayKey)
160 - }
161 - return nil
162 -}
66
164 -func (s *RelaySet) bootstrapRelayURLsLocked() []string {
165 - if len(s.knownRelayURLs) == 0 {
166 - return nil
167 - }
168 -
169 - out := make([]string, 0, len(s.knownRelayURLs))
170 - for _, relayURL := range s.knownRelayURLs {
171 - if s.isSelfRelayURLLocked(relayURL) || s.localByURL[relayURL].Banned {
172 - continue
173 - }
174 - out = append(out, relayURL)
175 - }
176 - if len(out) == 0 {
177 - return nil
178 - }
179 - return out
180 -}
67 + filtered := utils.RemoveRelayURL(normalized, s.selfRelayURL)
68
182 -func (s *RelaySet) descriptorProjectionsLocked() []relayDescriptorProjection {
183 - if len(s.relays) == 0 {
184 - return nil
69 + keep := make(map[string]struct{}, len(filtered))
70 + for _, relayURL := range filtered {
71 + keep[relayURL] = struct{}{}
72 }
73
187 - out := make([]relayDescriptorProjection, 0, len(s.relays))
188 - for _, record := range s.relays {
189 - if s.isSelfRelayDescriptorLocked(record.Descriptor) {
74 + for _, relayURL := range s.knownRelayURLs {
75 + if _, ok := keep[relayURL]; ok {
76 continue
77 }
192 - relayURL := strings.TrimSpace(record.Descriptor.APIHTTPSAddr)
193 - if relayURL == "" {
78 + state := s.relays[relayURL]
79 + if state.hasDescriptor() {
80 continue
81 }
196 - bootstrap := false
197 - if !s.isSelfRelayURLLocked(relayURL) {
198 - for _, candidate := range s.knownRelayURLs {
199 - if candidate == relayURL {
200 - bootstrap = true
201 - break
202 - }
203 - }
82 + if !s.policy.KeepState(state) {
83 + delete(s.relays, relayURL)
84 }
205 - local := s.localByURL[relayURL]
206 - record.Banned = local.Banned
207 - record.status = local.status
208 - record.consecutiveFailures = local.consecutiveFailures
209 - out = append(out, relayDescriptorProjection{
210 - state: record,
211 - relayURL: relayURL,
212 - bootstrap: bootstrap,
213 - })
214 - }
215 - if len(out) == 0 {
216 - return nil
85 }
218 - sort.Slice(out, func(i, j int) bool {
219 - return out[i].relayURL < out[j].relayURL
220 - })
221 - return out
222 -}
86
224 -func (s *RelaySet) ActiveRelayURLs() []string {
225 - s.mu.RLock()
226 - defer s.mu.RUnlock()
87 + s.knownRelayURLs = append([]string(nil), filtered...)
88 + return nil
89 +}
90
228 - return selectRelayURLsLocked(s.relayCandidatesLocked(time.Now().UTC()), SelectionOptions{})
91 +func (s *RelaySet) relayURLForKeyLocked(relayKey string) string {
92 + if relayKey == "" {
93 + return ""
94 + }
95 + for relayURL, state := range s.relays {
96 + if state.hasDescriptor() && state.Descriptor.Key() == relayKey {
97 + return relayURL
98 + }
99 + }
100 + return ""
101 }
102
231 -func (s *RelaySet) RelayCandidates() []RelayCandidate {
103 +func (s *RelaySet) RelayStates() []RelayState {
104 s.mu.RLock()
105 defer s.mu.RUnlock()
106
235 - return s.relayCandidatesLocked(time.Now().UTC())
107 + states := s.relayStatesLocked()
108 + for i := range states {
109 + states[i] = s.policy.Decide(states[i])
110 + }
111 + return states
112 }
113
238 -func (s *RelaySet) SelectRelayURLs(opts SelectionOptions) []string {
114 +func (s *RelaySet) ActiveRelays() []RelayState {
115 s.mu.RLock()
116 defer s.mu.RUnlock()
117
242 - return selectRelayURLsLocked(s.relayCandidatesLocked(time.Now().UTC()), opts)
243 -}
244 -
245 -func selectRelayURLsLocked(candidates []RelayCandidate, opts SelectionOptions) []string {
246 - if len(candidates) == 0 {
247 - return nil
248 - }
249 -
250 - limit := opts.Limit
251 - out := make([]string, 0, len(candidates))
252 - seen := make(map[string]struct{}, len(candidates))
253 - for _, candidate := range candidates {
254 - relayURL := relayCandidateURL(candidate)
255 - if relayURL == "" {
256 - continue
257 - }
258 - if _, ok := seen[relayURL]; ok {
259 - continue
260 - }
261 - seen[relayURL] = struct{}{}
262 - out = append(out, relayURL)
263 - if limit > 0 && len(out) >= limit {
264 - break
118 + states := s.relayStatesLocked()
119 + out := make([]RelayState, 0, len(states))
120 + for _, state := range states {
121 + state = s.policy.Decide(state)
122 + if state.Active && state.Descriptor.APIHTTPSAddr != "" {
123 + out = append(out, state)
124 }
125 }
126 if len(out) == 0 {
@@ -270,449 +129,167 @@ func selectRelayURLsLocked(candidates []RelayCandidate, opts SelectionOptions) [
129 return out
130 }
131
273 -func relayCandidateURL(candidate RelayCandidate) string {
274 - relayURL := strings.TrimSpace(candidate.Descriptor.APIHTTPSAddr)
275 - if relayURL != "" {
276 - return relayURL
277 - }
278 - return strings.TrimSpace(candidate.Descriptor.RelayID)
279 -}
132 +func (s *RelaySet) relayStatesLocked() []RelayState {
133 + out := make([]RelayState, 0, len(s.knownRelayURLs)+len(s.relays))
134 + seen := make(map[string]struct{}, len(s.knownRelayURLs)+len(s.relays))
135
281 -func (s *RelaySet) relayCandidatesLocked(now time.Time) []RelayCandidate {
282 - bootstrapRelayURLs := s.bootstrapRelayURLsLocked()
283 - projections := s.descriptorProjectionsLocked()
284 -
285 - out := make([]RelayCandidate, 0, len(bootstrapRelayURLs)+len(projections))
286 - seen := make(map[string]struct{}, len(bootstrapRelayURLs)+len(projections))
287 - for _, relayURL := range bootstrapRelayURLs {
136 + for _, relayURL := range s.knownRelayURLs {
137 + if relayURL == "" {
138 + continue
139 + }
140 if _, ok := seen[relayURL]; ok {
141 continue
142 }
143 seen[relayURL] = struct{}{}
144
293 - candidate := RelayCandidate{
294 - Bootstrap: true,
295 - Descriptor: types.RelayDescriptor{
296 - Identity: types.Identity{
297 - Name: utils.PortalRootHost(relayURL),
298 - },
299 - RelayID: relayURL,
300 - APIHTTPSAddr: relayURL,
301 - Version: 1,
302 - },
145 + state, ok := s.relays[relayURL]
146 + if !ok {
147 + state = newRelayHintState(relayURL)
148 }
304 - if relayKey, ok := s.relayKeysByURL[relayURL]; ok {
305 - if record, ok := s.relays[relayKey]; ok && record.Descriptor.APIHTTPSAddr != "" {
306 - candidate.Descriptor = record.Descriptor
307 - candidate.FirstSeenAt = record.FirstSeenAt
308 - candidate.LastSeenAt = record.LastSeenAt
309 - }
149 + if selfRelay(state, s.selfRelayKey, s.selfRelayURL) {
150 + continue
151 }
311 - out = append(out, candidate)
152 + state.Bootstrap = true
153 + out = append(out, state)
154 }
313 - for _, projection := range projections {
314 - if projection.state.Banned || projection.state.status != relayStatusConfirmed || relayExpiredAt(projection.state, now) {
155 +
156 + descriptorStates := make([]RelayState, 0, len(s.relays))
157 + for relayURL, record := range s.relays {
158 + if !record.hasDescriptor() {
159 continue
160 }
317 - if _, ok := seen[projection.relayURL]; ok {
161 + relayURL = strings.TrimSpace(record.Descriptor.APIHTTPSAddr)
162 + if relayURL == "" {
163 continue
164 }
320 - seen[projection.relayURL] = struct{}{}
321 - out = append(out, RelayCandidate{
322 - Descriptor: projection.state.Descriptor,
323 - Bootstrap: projection.bootstrap,
324 - FirstSeenAt: projection.state.FirstSeenAt,
325 - LastSeenAt: projection.state.LastSeenAt,
326 - })
327 - }
328 - if len(out) == 0 {
329 - return nil
330 - }
331 - return out
332 -}
333 -
334 -func (s *RelaySet) BootstrapDescriptors() []types.RelayDescriptor {
335 - s.mu.RLock()
336 - defer s.mu.RUnlock()
337 -
338 - bootstrapRelayURLs := s.bootstrapRelayURLsLocked()
339 - if len(bootstrapRelayURLs) == 0 {
340 - return nil
341 - }
342 -
343 - out := make([]types.RelayDescriptor, 0, len(bootstrapRelayURLs))
344 - for _, relayURL := range bootstrapRelayURLs {
345 - if relayKey, ok := s.relayKeysByURL[relayURL]; ok {
346 - if record, ok := s.relays[relayKey]; ok && record.Descriptor.APIHTTPSAddr != "" {
347 - out = append(out, record.Descriptor)
348 - continue
349 - }
350 - }
351 - out = append(out, types.RelayDescriptor{
352 - Identity: types.Identity{
353 - Name: utils.PortalRootHost(relayURL),
354 - },
355 - RelayID: relayURL,
356 - APIHTTPSAddr: relayURL,
357 - Version: 1,
358 - })
359 - }
360 - if len(out) == 0 {
361 - return nil
362 - }
363 - return out
364 -}
365 -
366 -func (s *RelaySet) BanRelayURL(relayURL string) {
367 - s.mu.Lock()
368 - defer s.mu.Unlock()
369 - relayURL = strings.TrimSpace(relayURL)
370 - if relayURL == "" {
371 - return
372 - }
373 -
374 - state := s.localByURL[relayURL]
375 - state.Banned = true
376 - s.storeLocalStateLocked(relayURL, state)
377 -}
378 -
379 -func (s *RelaySet) AdvertisedDescriptors() []types.RelayDescriptor {
380 - s.mu.RLock()
381 - defer s.mu.RUnlock()
382 -
383 - now := time.Now().UTC()
384 - projections := s.descriptorProjectionsLocked()
385 - if len(projections) == 0 {
386 - return nil
387 - }
388 -
389 - out := make([]types.RelayDescriptor, 0, len(projections))
390 - for _, projection := range projections {
391 - if projection.state.Banned || projection.state.status != relayStatusConfirmed || relayExpiredAt(projection.state, now) {
165 + if selfRelay(record, s.selfRelayKey, s.selfRelayURL) {
166 continue
167 }
394 - out = append(out, projection.state.Descriptor)
395 - }
396 - if len(out) == 0 {
397 - return nil
398 - }
399 - return out
400 -}
401 -
402 -func (s *RelaySet) confirmableDescriptors() []types.RelayDescriptor {
403 - s.mu.RLock()
404 - defer s.mu.RUnlock()
405 -
406 - now := time.Now().UTC()
407 - projections := s.descriptorProjectionsLocked()
408 - if len(projections) == 0 {
409 - return nil
410 - }
411 -
412 - out := make([]types.RelayDescriptor, 0, len(projections))
413 - for _, projection := range projections {
414 - if projection.bootstrap || projection.state.Banned || relayExpiredAt(projection.state, now) {
168 + if _, ok := seen[relayURL]; ok {
169 continue
170 }
417 - out = append(out, projection.state.Descriptor)
418 - }
419 - if len(out) == 0 {
420 - return nil
421 - }
422 - return out
423 -}
171
425 -func (s *RelaySet) SyncableDescriptors() []types.RelayDescriptor {
426 - s.mu.RLock()
427 - defer s.mu.RUnlock()
428 -
429 - now := time.Now().UTC()
430 - projections := s.descriptorProjectionsLocked()
431 - if len(projections) == 0 {
432 - return nil
172 + record.Bootstrap = slices.Contains(s.knownRelayURLs, relayURL)
173 + descriptorStates = append(descriptorStates, record)
174 }
175 + sort.Slice(descriptorStates, func(i, j int) bool {
176 + return descriptorStates[i].Descriptor.APIHTTPSAddr < descriptorStates[j].Descriptor.APIHTTPSAddr
177 + })
178 + out = append(out, descriptorStates...)
179
435 - out := make([]types.RelayDescriptor, 0, len(projections))
436 - for _, projection := range projections {
437 - if projection.bootstrap || projection.state.Banned || relayExpiredAt(projection.state, now) || !projection.state.Descriptor.SupportsOverlayPeer {
438 - continue
439 - }
440 - out = append(out, projection.state.Descriptor)
441 - }
180 if len(out) == 0 {
181 return nil
182 }
183 return out
184 }
185
448 -func (s *RelaySet) View() map[string]RelayState {
186 +func (s *RelaySet) AdvertisedDescriptors() []types.RelayDescriptor {
187 s.mu.RLock()
188 defer s.mu.RUnlock()
189
452 - now := time.Now().UTC()
453 - projections := s.descriptorProjectionsLocked()
454 - if len(projections) == 0 {
455 - return nil
456 - }
457 -
458 - view := make(map[string]RelayState, len(projections))
459 - for _, projection := range projections {
460 - relayKey := projection.state.Descriptor.Key()
461 - if relayKey == "" {
462 - continue
463 - }
464 - expired := relayExpiredAt(projection.state, now)
465 - view[relayKey] = RelayState{
466 - Descriptor: projection.state.Descriptor,
467 - FirstSeenAt: projection.state.FirstSeenAt,
468 - LastSeenAt: projection.state.LastSeenAt,
469 - Banned: projection.state.Banned,
470 - Expired: expired,
471 - }
472 - }
473 - if len(view) == 0 {
474 - return nil
475 - }
476 - return view
190 + return s.policy.AdvertisedDescriptors(s.relayStatesLocked())
191 }
192
479 -func (s *RelaySet) SetBootstrapRelayURLs(inputs []string) error {
480 - normalized, err := utils.NormalizeRelayURLs(inputs...)
481 - if err != nil {
482 - return err
483 - }
484 -
193 +func (s *RelaySet) BanRelayURL(relayURL string) {
194 s.mu.Lock()
195 defer s.mu.Unlock()
487 -
488 - filtered := utils.RemoveRelayURL(normalized, s.selfRelayURL)
489 -
490 - keep := make(map[string]struct{}, len(filtered))
491 - for _, relayURL := range filtered {
492 - keep[relayURL] = struct{}{}
196 + relayURL = strings.TrimSpace(relayURL)
197 + if relayURL == "" {
198 + return
199 }
200
495 - for _, relayURL := range s.knownRelayURLs {
496 - if _, ok := keep[relayURL]; ok {
497 - continue
498 - }
499 - if _, ok := s.relayKeysByURL[relayURL]; ok {
500 - continue
501 - }
502 - if state := s.localByURL[relayURL]; state.isDefaultLocalState() {
503 - delete(s.localByURL, relayURL)
504 - }
201 + state := s.relays[relayURL]
202 + if state.Descriptor.APIHTTPSAddr == "" {
203 + state = newRelayHintState(relayURL)
204 }
506 -
507 - s.knownRelayURLs = append([]string(nil), filtered...)
508 - return nil
205 + state = s.policy.OnBanned(state)
206 + s.relays[relayURL] = state
207 }
208
511 -func (s *RelaySet) storeDescriptorLocked(desc types.RelayDescriptor, now time.Time) (string, bool, bool, error) {
209 +func (s *RelaySet) storeDescriptorLocked(state RelayState, policy RelayPolicy) (bool, bool, error) {
210 + desc := state.Descriptor
211 relayKey := desc.Key()
212 if relayKey == "" {
514 - return "", false, false, errors.New("descriptor identity is required")
213 + return false, false, errors.New("descriptor identity is required")
214 }
516 - if knownRelayKey, ok := s.relayKeysByURL[desc.APIHTTPSAddr]; ok && knownRelayKey != relayKey {
517 - return "", false, false, errors.New("descriptor identity does not match known relay url")
518 - }
519 - if now.IsZero() {
520 - now = time.Now().UTC()
215 + if existing := s.relays[desc.APIHTTPSAddr]; existing.hasDescriptor() && existing.Descriptor.Key() != relayKey {
216 + return false, false, errors.New("descriptor identity does not match known relay url")
217 }
218
523 - record, ok := s.relays[relayKey]
524 - added := !ok
525 - if !ok {
526 - record.FirstSeenAt = now
527 - }
528 - previousURL := record.Descriptor.APIHTTPSAddr
219 + previousURL := s.relayURLForKeyLocked(relayKey)
220 + record := s.relays[desc.APIHTTPSAddr]
221 previousDescriptor := record.Descriptor
530 - record.Descriptor = desc
531 - record.LastSeenAt = now
532 - s.relays[relayKey] = record
533 - s.relayKeysByURL[desc.APIHTTPSAddr] = relayKey
222 if previousURL != "" && previousURL != desc.APIHTTPSAddr {
535 - delete(s.relayKeysByURL, previousURL)
536 - state := s.localByURL[previousURL]
537 - bootstrap := false
538 - if !s.isSelfRelayURLLocked(previousURL) {
539 - if slices.Contains(s.knownRelayURLs, previousURL) {
540 - bootstrap = true
541 - }
223 + previous := s.relays[previousURL]
224 + previousDescriptor = previous.Descriptor
225 + if !policy.KeepState(record) {
226 + record.Status = previous.Status
227 + record.consecutiveFailures = previous.consecutiveFailures
228 }
543 - if !bootstrap && state.isDefaultLocalState() {
544 - delete(s.localByURL, previousURL)
229 + delete(s.relays, previousURL)
230 + }
231 +
232 + added := previousURL == ""
233 + if !policy.KeepState(record) {
234 + record = state
235 + } else {
236 + if record.FirstSeenAt.IsZero() {
237 + record.FirstSeenAt = state.FirstSeenAt
238 }
239 + record.Descriptor = desc
240 + record.LastSeenAt = state.LastSeenAt
241 }
242 + s.relays[desc.APIHTTPSAddr] = record
243
244 changed := added || !reflect.DeepEqual(previousDescriptor, desc)
549 - return relayKey, added, changed, nil
245 + return added, changed, nil
246 }
247
552 -func (s *RelaySet) applyDiscoveryResponse(targetIdentity types.Identity, targetURL string, resp types.DiscoveryResponse, now time.Time, requireOverlay bool) (relaySetChanged bool, warnErr error, err error) {
248 +func (s *RelaySet) ApplyRelayDiscoveryResponse(targetIdentity types.Identity, targetURL string, resp types.DiscoveryResponse, now time.Time) (relaySetChanged bool, err error) {
249 if now.IsZero() {
250 now = time.Now().UTC()
251 } else {
252 now = now.UTC()
253 }
254
559 - protocolVersion := strings.TrimSpace(resp.ProtocolVersion)
560 - if protocolVersion != types.ProtocolVersion {
561 - err := fmt.Errorf("relay protocol version mismatch: relay=%q client=%q", protocolVersion, types.ProtocolVersion)
562 - return false, err, err
563 - }
564 -
565 - normalize := func(desc types.RelayDescriptor) (types.RelayDescriptor, error) {
566 - normalized, err := utils.NormalizeDescriptor(desc)
567 - if err != nil {
568 - return types.RelayDescriptor{}, err
569 - }
255 + s.mu.Lock()
256 + defer s.mu.Unlock()
257
571 - switch {
572 - case normalized.Name == "":
573 - return types.RelayDescriptor{}, errors.New("identity.name is required")
574 - case normalized.APIHTTPSAddr == "":
575 - return types.RelayDescriptor{}, errors.New("api_https_addr is required")
576 - case normalized.RelayID == "":
577 - return types.RelayDescriptor{}, errors.New("relay_id is required")
578 - case normalized.APIHTTPSAddr != "" && normalized.RelayID != normalized.APIHTTPSAddr:
579 - return types.RelayDescriptor{}, errors.New("relay_id must match api_https_addr")
580 - case normalized.Sequence == 0:
581 - return types.RelayDescriptor{}, errors.New("sequence is required")
582 - case normalized.Version == 0:
583 - return types.RelayDescriptor{}, errors.New("version is required")
584 - case normalized.IssuedAt.IsZero():
585 - return types.RelayDescriptor{}, errors.New("issued_at is required")
586 - case normalized.ExpiresAt.IsZero():
587 - return types.RelayDescriptor{}, errors.New("expires_at is required")
588 - case normalized.ExpiresAt.Before(now):
589 - return types.RelayDescriptor{}, errors.New("descriptor expired")
590 - case normalized.IssuedAt.After(normalized.ExpiresAt):
591 - return types.RelayDescriptor{}, errors.New("issued_at must be before expires_at")
592 - }
593 - return normalized, nil
594 - }
258 + policy := s.policy
259
596 - selfDescriptor, err := normalize(resp.Self)
260 + selfState, relayStates, err := policy.DiscoveryStates(targetIdentity, targetURL, s.selfRelayKey, s.selfRelayURL, resp, now)
261 if err != nil {
598 - return false, err, err
599 - }
600 - if requireOverlay {
601 - if err := RequireOverlayRelayDescriptor(selfDescriptor); err != nil {
602 - return false, warnErr, err
603 - }
262 + return false, err
263 }
264
606 - seen := map[string]struct{}{selfDescriptor.Key(): {}}
607 - relayDescriptors := make([]types.RelayDescriptor, 0, len(resp.Relays))
608 - for _, descriptor := range resp.Relays {
609 - relayDescriptor, err := normalize(descriptor)
610 - if err != nil {
611 - log.Warn().
612 - Err(err).
613 - Str("relay", strings.TrimSpace(descriptor.APIHTTPSAddr)).
614 - Str("name", strings.TrimSpace(descriptor.Name)).
615 - Msg("skipping invalid discovery relay hint")
616 - continue
617 - }
618 - relayKey := relayDescriptor.Key()
619 - if _, ok := seen[relayKey]; ok {
620 - log.Debug().
621 - Str("relay", relayDescriptor.APIHTTPSAddr).
622 - Str("identity_key", relayKey).
623 - Msg("skipping duplicate discovery relay hint")
624 - continue
625 - }
626 - seen[relayKey] = struct{}{}
627 - relayDescriptors = append(relayDescriptors, relayDescriptor)
628 - }
629 -
630 - if strings.TrimSpace(targetIdentity.Name) == "" && strings.TrimSpace(targetIdentity.Address) == "" {
631 - return false, warnErr, errors.New("target relay identity is required")
632 - }
633 - targetName := strings.TrimSpace(targetIdentity.Name)
634 - if targetName != "" {
635 - normalizedTargetName := utils.NormalizeHostname(targetName)
636 - if selfDescriptor.Name != normalizedTargetName {
637 - return false, warnErr, errors.New("descriptor name does not match target relay")
638 - }
639 - }
640 - targetAddress := strings.TrimSpace(targetIdentity.Address)
641 - if targetAddress != "" {
642 - normalizedTargetAddress, err := utils.NormalizeEVMAddress(targetAddress)
643 - if err != nil {
644 - return false, warnErr, err
645 - }
646 - if selfDescriptor.Address != normalizedTargetAddress {
647 - return false, warnErr, errors.New("descriptor address does not match target relay")
648 - }
649 - }
650 - if targetURL != "" {
651 - normalizedTargetURL, err := utils.NormalizeRelayURL(targetURL)
652 - if err != nil {
653 - return false, warnErr, err
654 - }
655 - if selfDescriptor.APIHTTPSAddr != normalizedTargetURL {
656 - return false, warnErr, errors.New("descriptor api_https_addr does not match target url")
657 - }
658 - }
659 -
660 - s.mu.Lock()
661 - defer s.mu.Unlock()
662 -
663 - apply := func(desc types.RelayDescriptor, advertise bool) error {
664 - if !advertise && s.isSelfRelayDescriptorLocked(desc) {
665 - return nil
666 - }
667 - if !advertise && requireOverlay {
668 - if err := RequireOverlayRelayDescriptor(desc); err != nil {
669 - if warnErr == nil {
670 - warnErr = err
671 - }
672 - return nil
673 - }
674 - }
675 -
676 - _, added, descriptorChanged, err := s.storeDescriptorLocked(desc, now)
265 + apply := func(state RelayState, advertise bool) error {
266 + desc := state.Descriptor
267 + added, descriptorChanged, err := s.storeDescriptorLocked(state, policy)
268 if err != nil {
269 return err
270 }
271
681 - localState := s.localByURL[desc.APIHTTPSAddr]
682 - previousState := localState
683 - if advertise {
684 - localState.status = relayStatusConfirmed
685 - localState.consecutiveFailures = 0
686 - } else if localState.status != relayStatusConfirmed {
687 - localState.status = relayStatusHinted
688 - localState.consecutiveFailures = 0
689 - }
690 - s.storeLocalStateLocked(desc.APIHTTPSAddr, localState)
272 + storedState := s.relays[desc.APIHTTPSAddr]
273 + previousState := storedState
274 + storedState = policy.OnDiscovered(storedState, advertise)
275 + s.relays[desc.APIHTTPSAddr] = storedState
276
692 - changed := added || descriptorChanged || !reflect.DeepEqual(previousState, localState)
277 + changed := added || descriptorChanged || !reflect.DeepEqual(previousState, s.relays[desc.APIHTTPSAddr])
278 if changed {
279 relaySetChanged = true
280 }
281 return nil
282 }
283
699 - if err := apply(selfDescriptor, true); err != nil {
700 - return false, warnErr, err
284 + if err := apply(selfState, true); err != nil {
285 + return false, err
286 }
702 - for _, relayDescriptor := range relayDescriptors {
703 - if err := apply(relayDescriptor, false); err != nil {
704 - return false, warnErr, err
287 + for _, relayState := range relayStates {
288 + if err := apply(relayState, false); err != nil {
289 + return false, err
290 }
291 }
707 - return relaySetChanged, warnErr, nil
708 -}
709 -
710 -func (s *RelaySet) ApplyRelayDiscoveryResponse(targetIdentity types.Identity, targetURL string, resp types.DiscoveryResponse, now time.Time) (relaySetChanged bool, warnErr error, err error) {
711 - return s.applyDiscoveryResponse(targetIdentity, targetURL, resp, now, false)
712 -}
713 -
714 -func (s *RelaySet) ApplyOverlayRelayDiscoveryResponse(targetIdentity types.Identity, targetURL string, resp types.DiscoveryResponse, now time.Time) (relaySetChanged bool, warnErr error, err error) {
715 - return s.applyDiscoveryResponse(targetIdentity, targetURL, resp, now, true)
292 + return relaySetChanged, nil
293 }
294
295 func (s *RelaySet) RecordDiscoveryFailure(identity types.Identity, relayURL string, err error, recoveryFailures int) (expired bool, expireReason string, consecutiveFailures int) {
@@ -725,36 +302,18 @@ func (s *RelaySet) RecordDiscoveryFailure(identity types.Identity, relayURL stri
302 s.mu.Lock()
303 defer s.mu.Unlock()
304
728 - record, ok := s.relays[relayKey]
729 - if !ok {
730 - return false, "", 0
731 - }
732 - if relayURL == "" || s.relayKeysByURL[relayURL] != relayKey {
733 - relayURL = record.Descriptor.APIHTTPSAddr
305 + if relayURL == "" || s.relays[relayURL].Descriptor.Key() != relayKey {
306 + relayURL = s.relayURLForKeyLocked(relayKey)
307 }
308 if relayURL == "" {
309 return false, "", 0
310 }
311
739 - localState := s.localByURL[relayURL]
740 - localState.consecutiveFailures++
741 - s.storeLocalStateLocked(relayURL, localState)
742 - if localState.status != relayStatusExpired && localState.consecutiveFailures >= recoveryFailures {
743 - state := s.localByURL[record.Descriptor.APIHTTPSAddr]
744 - state.status = relayStatusExpired
745 - s.storeLocalStateLocked(record.Descriptor.APIHTTPSAddr, state)
746 - return true, "recovery", localState.consecutiveFailures
747 - }
748 -
749 - var apiErr *types.APIRequestError
750 - if errors.As(err, &apiErr) &&
751 - (apiErr.StatusCode == http.StatusForbidden ||
752 - apiErr.StatusCode == http.StatusNotFound ||
753 - apiErr.StatusCode == http.StatusGone) {
754 - state := s.localByURL[record.Descriptor.APIHTTPSAddr]
755 - state.status = relayStatusExpired
756 - s.storeLocalStateLocked(record.Descriptor.APIHTTPSAddr, state)
757 - return true, "status", localState.consecutiveFailures
312 + state, ok := s.relays[relayURL]
313 + if !ok {
314 + return false, "", 0
315 }
759 - return false, "", localState.consecutiveFailures
316 + state, expired, expireReason = s.policy.OnFailure(state, err, recoveryFailures)
317 + s.relays[relayURL] = state
318 + return expired, expireReason, state.consecutiveFailures
319 }
portal/overlay/overlay.go
+8 -10
@@ -209,24 +209,22 @@ func (o *Overlay) DiscoverRelay(ctx context.Context, relay types.RelayDescriptor
209 return resp, nil
210 }
211
212 -func (o *Overlay) Sync(view map[string]discovery.RelayState) error {
212 +func (o *Overlay) Sync(relays []discovery.RelayState) error {
213 if o == nil || o.stack == nil {
214 return nil
215 }
216 - return o.stack.ApplyPeers(peersForView(o.cfg.PublicKey, view))
216 + return o.stack.ApplyPeers(peersForRelays(o.cfg.PublicKey, relays))
217 }
218
219 -func peersForView(publicKey string, view map[string]discovery.RelayState) []desiredPeer {
220 - peers := make([]desiredPeer, 0, len(view))
221 - for _, relay := range view {
222 - if relay.Expired || relay.Banned {
219 +func peersForRelays(publicKey string, relays []discovery.RelayState) []desiredPeer {
220 + peers := make([]desiredPeer, 0, len(relays))
221 + for _, relay := range relays {
222 + if !relay.OverlayPeer {
223 continue
224 }
225 +
226 desc := relay.Descriptor
226 - if desc.WireGuardPublicKey == publicKey || !desc.SupportsOverlayPeer {
227 - continue
228 - }
229 - if desc.WireGuardPublicKey == "" || desc.WireGuardEndpoint == "" || desc.OverlayIPv4 == "" {
227 + if desc.WireGuardPublicKey == publicKey {
228 continue
229 }
230
portal/server.go
+1 -1
@@ -614,7 +614,7 @@ func (s *Server) startOverlay() (*overlay.Overlay, error) {
614 return nil, fmt.Errorf("start wireguard overlay: %w", err)
615 }
616
617 - if err := overlay.Sync(s.relaySet.View()); err != nil {
617 + if err := overlay.Sync(s.relaySet.RelayStates()); err != nil {
618 _ = overlay.Shutdown(context.Background())
619 return nil, fmt.Errorf("sync wireguard peers: %w", err)
620 }
portal/server_test.go
+20 -8
@@ -51,11 +51,8 @@ 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 - _, warnErr, err := set.ApplyRelayDiscoveryResponse(identity, targetURL, resp, now)
55 - if err != nil {
56 - return err
57 - }
58 - return warnErr
54 + _, err := set.ApplyRelayDiscoveryResponse(identity, targetURL, resp, now)
55 + return err
56 }
57
58 func tempIdentityPath(t *testing.T) string {
@@ -511,7 +508,12 @@ func TestServerSetBootstrapRelayURLsAllowsLoopbackButSkipsSelfRelay(t *testing.T
508 t.Fatalf("SetBootstrapRelayURLs() error = %v", err)
509 }
510 advertisedDescriptors := server.relaySet.AdvertisedDescriptors()
514 - knownURLs := append([]string(nil), server.relaySet.ActiveRelayURLs()...)
511 + knownURLs := make([]string, 0)
512 + for _, state := range server.relaySet.RelayStates() {
513 + if state.Active && state.Descriptor.APIHTTPSAddr != "" {
514 + knownURLs = append(knownURLs, state.Descriptor.APIHTTPSAddr)
515 + }
516 + }
517 sort.Strings(knownURLs)
518 if !reflect.DeepEqual(knownURLs, []string{
519 "https://bootstrap.example.com",
@@ -564,7 +566,12 @@ func TestServerDiscoverySkipsSelfRelayHint(t *testing.T) {
566 t.Fatalf("ApplyRelayDiscoveryResponse() error = %v", err)
567 }
568
567 - knownURLs := append([]string(nil), server.relaySet.ActiveRelayURLs()...)
569 + knownURLs := make([]string, 0)
570 + for _, state := range server.relaySet.RelayStates() {
571 + if state.Active && state.Descriptor.APIHTTPSAddr != "" {
572 + knownURLs = append(knownURLs, state.Descriptor.APIHTTPSAddr)
573 + }
574 + }
575 sort.Strings(knownURLs)
576 if !reflect.DeepEqual(knownURLs, []string{"https://bootstrap.example.com"}) {
577 t.Fatalf("ActiveRelayURLs() = %v, want self hint excluded", knownURLs)
@@ -609,7 +616,12 @@ func TestServerRecordVerifiedDiscoveryPeerRequiresDirectConfirmation(t *testing.
616 if err != nil {
617 t.Fatalf("applyRelayDiscoveryResponse() hinted error = %v", err)
618 }
612 - knownURLs := append([]string(nil), server.relaySet.ActiveRelayURLs()...)
619 + knownURLs := make([]string, 0)
620 + for _, state := range server.relaySet.RelayStates() {
621 + if state.Active && state.Descriptor.APIHTTPSAddr != "" {
622 + knownURLs = append(knownURLs, state.Descriptor.APIHTTPSAddr)
623 + }
624 + }
625 sort.Strings(knownURLs)
626 if !reflect.DeepEqual(knownURLs, []string{"https://bootstrap.example.com"}) {
627 t.Fatalf("ActiveRelayURLs() = %v, want [%q]", knownURLs, "https://bootstrap.example.com")
sdk/expose.go
+11 -5
@@ -159,7 +159,11 @@ func (e *Exposure) runDiscoveryLoop(ctx context.Context) {
159 }
160
161 func (e *Exposure) ActiveRelayURLs() []string {
162 - return e.relaySet.ActiveRelayURLs()
162 + var out []string
163 + for _, state := range e.relaySet.ActiveRelays() {
164 + out = append(out, state.Descriptor.APIHTTPSAddr)
165 + }
166 + return out
167 }
168
169 func (e *Exposure) Addr() net.Addr {
@@ -254,11 +258,10 @@ func (e *Exposure) WaitDatagramReady(ctx context.Context) ([]string, error) {
258
259 func (e *Exposure) RunHTTP(ctx context.Context, handler http.Handler, localAddr string) error {
260 var relayListener net.Listener
257 - activeRelayURLs := e.relaySet.ActiveRelayURLs()
261 e.listenerMu.RLock()
262 activeListeners := make([]*Listener, 0, len(e.relayListeners))
260 - for _, relayURL := range activeRelayURLs {
261 - listener, ok := e.relayListeners[relayURL]
263 + for _, state := range e.relaySet.ActiveRelays() {
264 + listener, ok := e.relayListeners[state.Descriptor.APIHTTPSAddr]
265 if !ok {
266 continue
267 }
@@ -365,12 +368,15 @@ func (e *Exposure) Close() error {
368 }
369
370 func (e *Exposure) reconcileRelayListeners(failOnError bool) error {
368 - activeRelayURLs := e.relaySet.ActiveRelayURLs()
371 e.listenerMu.Lock()
372 currentRelayURLs := make([]string, 0, len(e.relayListeners))
373 for relayURL := range e.relayListeners {
374 currentRelayURLs = append(currentRelayURLs, relayURL)
375 }
376 + activeRelayURLs := make([]string, 0, len(e.relayListeners))
377 + for _, state := range e.relaySet.ActiveRelays() {
378 + activeRelayURLs = append(activeRelayURLs, state.Descriptor.APIHTTPSAddr)
379 + }
380 missingRelayURLs := utils.FilterRelayURLs(activeRelayURLs, currentRelayURLs)
381 staleRelayURLs := utils.FilterRelayURLs(currentRelayURLs, activeRelayURLs)
382 staleListeners := make([]*Listener, 0, len(staleRelayURLs))
sdk/expose_test.go
+5 -8
@@ -43,11 +43,8 @@ func mustRelayDescriptor(t *testing.T, relayName, relayURL string) types.RelayDe
43
44 func applyRelayDiscovery(t *testing.T, set *discovery.RelaySet, identity types.Identity, targetURL string, resp types.DiscoveryResponse, now time.Time) error {
45 t.Helper()
46 - _, warnErr, err := set.ApplyRelayDiscoveryResponse(identity, targetURL, resp, now)
47 - if err != nil {
48 - return err
49 - }
50 - return warnErr
46 + _, err := set.ApplyRelayDiscoveryResponse(identity, targetURL, resp, now)
47 + return err
48 }
49
50 func TestExposureBanRelayURLMovesRelay(t *testing.T) {
@@ -83,7 +80,7 @@ func TestExposureBanRelayURLMovesRelay(t *testing.T) {
80 t.Fatalf("ActiveRelayURLs() = %v, want [%q]", got, relayB)
81 }
82
86 - knownRelayURLs := exposure.relaySet.ActiveRelayURLs()
83 + knownRelayURLs := exposure.ActiveRelayURLs()
84 exposure.listenerMu.RLock()
85 _, listenerExists := exposure.relayListeners[relayA]
86 exposure.listenerMu.RUnlock()
@@ -119,7 +116,7 @@ func TestExposureSetRelayURLsSkipsBannedRelay(t *testing.T) {
116 if got := exposure.ActiveRelayURLs(); len(got) != 1 || got[0] != relayA {
117 t.Fatalf("ActiveRelayURLs() = %v, want [%q]", got, relayA)
118 }
122 - knownRelayURLs := exposure.relaySet.ActiveRelayURLs()
119 + knownRelayURLs := exposure.ActiveRelayURLs()
120 if len(knownRelayURLs) != 1 || knownRelayURLs[0] != relayA {
121 t.Fatalf("knownRelayURLs = %v, want [%q]", knownRelayURLs, relayA)
122 }
@@ -169,7 +166,7 @@ func TestExposureSetRelayURLsRemovesStaleListener(t *testing.T) {
166 t.Fatal("stale relay listener was not closed")
167 }
168
172 - knownRelayURLs := exposure.relaySet.ActiveRelayURLs()
169 + knownRelayURLs := exposure.ActiveRelayURLs()
170 exposure.listenerMu.RLock()
171 _, relayAExists := exposure.relayListeners[relayA]
172 _, relayBExists := exposure.relayListeners[relayB]
sdk/mitm_test.go
+8 -3
@@ -228,8 +228,8 @@ func TestMITMProbeDetectionBansListener(t *testing.T) {
228 Reason: types.MITMProbeReasonExporterMismatch,
229 }, nil)
230
231 - for _, activeRelayURL := range listener.relaySet.ActiveRelayURLs() {
232 - if activeRelayURL == relayURL.String() {
231 + for _, state := range listener.relaySet.RelayStates() {
232 + if state.Active && state.Descriptor.APIHTTPSAddr == relayURL.String() {
233 t.Fatal("relay still active after mitm detection")
234 }
235 }
@@ -260,7 +260,12 @@ func TestMITMProbeDetectionWarnsWithoutBanningListener(t *testing.T) {
260 Reason: types.MITMProbeReasonExporterMismatch,
261 }, nil)
262
263 - activeRelayURLs := listener.relaySet.ActiveRelayURLs()
263 + activeRelayURLs := make([]string, 0)
264 + for _, state := range listener.relaySet.RelayStates() {
265 + if state.Active && state.Descriptor.APIHTTPSAddr != "" {
266 + activeRelayURLs = append(activeRelayURLs, state.Descriptor.APIHTTPSAddr)
267 + }
268 + }
269 if len(activeRelayURLs) != 1 || activeRelayURLs[0] != relayURL.String() {
270 t.Fatalf("ActiveRelayURLs() = %v, want [%q]", activeRelayURLs, relayURL.String())
271 }