refact: simplify discovery status and policy

rabbitprincess committed Apr 12, 2026 at 18:26 UTC 4895cdfca7d293fd916be1ad2d7d96f479890149
9 files changed +365 -431
portal/discovery/policy.go
+49 -232
@@ -2,284 +2,101 @@ package discovery
2
3 import (
4 "errors"
5 - "fmt"
5 "net/http"
7 - "sort"
8 - "strings"
6 "time"
7
8 "github.com/gosuda/portal-tunnel/v2/types"
12 - "github.com/gosuda/portal-tunnel/v2/utils"
9 )
10
15 -type RelayStatus uint8
16 -
17 -const (
18 - RelayStatusUnknown RelayStatus = iota
19 - RelayStatusHinted
20 - RelayStatusConfirmed
21 - RelayStatusInvalid
22 - RelayStatusExpired
23 - RelayStatusBanned
24 -)
25 -
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 - }
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 - }
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
11 +type RelayPolicy interface {
12 + SelectActive([]RelayState) []RelayState
13 + SelectAdvertised([]RelayState) []RelayState
14 + OnConfirmed(RelayState) RelayState
15 + OnHinted(RelayState) RelayState
16 + OnFailure(RelayState, error, int) (RelayState, bool, string)
17 + OnBanned(RelayState) RelayState
18 }
19
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 - }
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 - }
20 +type DefaultRelayPolicy struct{}
21
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) {
22 +func (p DefaultRelayPolicy) selectStates(states []RelayState, keep func(RelayState) bool) []RelayState {
23 + now := time.Now().UTC()
24 + out := make([]RelayState, 0, len(states))
25 + for _, state := range states {
26 + if !state.discoverable(now) {
27 continue
28 }
124 - relayKey := relayState.Descriptor.Key()
125 - if _, ok := seen[relayKey]; ok {
29 + if !keep(state) {
30 continue
31 }
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 - }
32 + out = append(out, state)
33 }
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 - }
34 + if len(out) == 0 {
35 + return nil
36 }
171 - return nil
37 + return out
38 }
39
174 -func (state RelayState) hasDescriptor() bool {
175 - return !state.LastSeenAt.IsZero() && state.Descriptor.Key() != "" && state.Descriptor.APIHTTPSAddr != ""
40 +func (p DefaultRelayPolicy) SelectActive(states []RelayState) []RelayState {
41 + return p.selectStates(states, func(state RelayState) bool {
42 + return state.Bootstrap || state.Confirmed
43 + })
44 }
45
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
46 +func (p DefaultRelayPolicy) SelectAdvertised(states []RelayState) []RelayState {
47 + return p.selectStates(states, func(state RelayState) bool {
48 + return state.Confirmed
49 + })
50 }
51
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
52 +func (DefaultRelayPolicy) OnConfirmed(state RelayState) RelayState {
53 + if state.Banned {
54 + return state
55 }
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
56 + state.Reachable = true
57 + state.Confirmed = true
58 + state.consecutiveFailures = 0
59 return state
60 }
61
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
62 +func (DefaultRelayPolicy) OnHinted(state RelayState) RelayState {
63 + if state.Banned {
64 return state
65 }
226 - if state.Status != RelayStatusConfirmed {
227 - state.Status = RelayStatusHinted
66 + state.Reachable = true
67 + if !state.Confirmed {
68 state.consecutiveFailures = 0
69 }
70 return state
71 }
72
73 func (DefaultRelayPolicy) OnFailure(state RelayState, err error, recoveryFailures int) (RelayState, bool, string) {
234 - if state.Status == RelayStatusBanned {
74 + if state.Banned {
75 return state, false, ""
76 }
77 state.consecutiveFailures++
238 - if state.Status != RelayStatusExpired && state.consecutiveFailures >= recoveryFailures {
239 - state.Status = RelayStatusExpired
240 - return state, true, "recovery"
78 + expire := func(reason string) (RelayState, bool, string) {
79 + if !state.Reachable {
80 + return state, false, ""
81 + }
82 + state.Reachable = false
83 + state.Confirmed = false
84 + return state, true, reason
85 }
86 var apiErr *types.APIRequestError
87 if errors.As(err, &apiErr) &&
88 (apiErr.StatusCode == http.StatusForbidden ||
89 apiErr.StatusCode == http.StatusNotFound ||
90 apiErr.StatusCode == http.StatusGone) {
247 - state.Status = RelayStatusExpired
248 - return state, true, "status"
91 + return expire("status")
92 + }
93 + if state.consecutiveFailures >= recoveryFailures {
94 + return expire("recovery")
95 }
96 return state, false, ""
97 }
98
99 func (DefaultRelayPolicy) OnBanned(state RelayState) RelayState {
254 - state.Status = RelayStatusBanned
100 + state.Banned = true
101 return state
102 }
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
+23 -11
@@ -16,7 +16,7 @@ import (
16 )
17
18 const (
19 - DiscoveryPollInterval = 1 * time.Minute
19 + DiscoveryPollInterval = 1 * time.Minute
20 defaultRecoveryFailures = 3
21 )
22
@@ -68,7 +68,7 @@ func (r *Refresher) Refresh(ctx context.Context) error {
68 if r.overlay == nil {
69 return ctx.Err()
70 }
71 - if err := r.overlay.Sync(r.relaySet.RelayStates()); err != nil {
71 + if err := r.overlay.Sync(r.relaySet.OverlayPeerStates()); err != nil {
72 log.Warn().
73 Err(err).
74 Msg("sync wireguard peers")
@@ -78,9 +78,13 @@ func (r *Refresher) Refresh(ctx context.Context) error {
78 }
79
80 func (r *Refresher) refreshHTTPS(ctx context.Context) error {
81 - states := r.relaySet.RelayStates()
81 + r.relaySet.mu.RLock()
82 + states := r.relaySet.relayStatesLocked()
83 + r.relaySet.mu.RUnlock()
84 +
85 + now := time.Now().UTC()
86 for _, state := range states {
83 - if state.Descriptor.APIHTTPSAddr == "" || !state.BootstrapDiscovery {
87 + if !state.discoverable(now) || !state.Bootstrap {
88 continue
89 }
90 relay := state.Descriptor
@@ -102,9 +106,13 @@ func (r *Refresher) refreshHTTPS(ctx context.Context) error {
106 return err
107 }
108
105 - states = r.relaySet.RelayStates()
109 + r.relaySet.mu.RLock()
110 + states = r.relaySet.relayStatesLocked()
111 + r.relaySet.mu.RUnlock()
112 +
113 + now = time.Now().UTC()
114 for _, state := range states {
107 - if state.Descriptor.APIHTTPSAddr == "" || !state.DirectDiscovery {
115 + if !state.discoverable(now) || state.Bootstrap {
116 continue
117 }
118 if r.overlay != nil && state.Descriptor.SupportsOverlayPeer {
@@ -144,11 +152,15 @@ func (r *Refresher) discoverHTTPS(ctx context.Context, relay types.RelayDescript
152 }
153
154 func (r *Refresher) refreshOverlay(ctx context.Context) error {
147 - for _, state := range r.relaySet.RelayStates() {
148 - if state.Descriptor.APIHTTPSAddr == "" || !state.OverlayDiscovery {
155 + r.relaySet.mu.RLock()
156 + states := r.relaySet.relayStatesLocked()
157 + r.relaySet.mu.RUnlock()
158 +
159 + now := time.Now().UTC()
160 + for _, state := range states {
161 + if !state.discoverable(now) || state.Bootstrap || !state.Descriptor.SupportsOverlayPeer {
162 continue
163 }
151 -
164 relay := state.Descriptor
165 var failureErr error
166
@@ -162,7 +174,7 @@ func (r *Refresher) refreshOverlay(ctx context.Context) error {
174 now := time.Now().UTC()
175 relaySetChanged, err := r.relaySet.ApplyRelayDiscoveryResponse(relay.Identity, relay.APIHTTPSAddr, resp, now)
176 if relaySetChanged {
165 - if syncErr := r.overlay.Sync(r.relaySet.RelayStates()); syncErr != nil {
177 + if syncErr := r.overlay.Sync(r.relaySet.OverlayPeerStates()); syncErr != nil {
178 log.Warn().
179 Err(syncErr).
180 Str("relay", relay.APIHTTPSAddr).
@@ -178,7 +190,7 @@ func (r *Refresher) refreshOverlay(ctx context.Context) error {
190
191 expired, expireReason, consecutiveFailures := r.relaySet.RecordDiscoveryFailure(relay.Identity, relay.APIHTTPSAddr, failureErr, r.overlayRecoveryFailures)
192 if expired {
181 - if syncErr := r.overlay.Sync(r.relaySet.RelayStates()); syncErr != nil && failureErr == nil {
193 + if syncErr := r.overlay.Sync(r.relaySet.OverlayPeerStates()); syncErr != nil && failureErr == nil {
194 failureErr = syncErr
195 }
196 }
portal/discovery/relayset.go
+174 -168
@@ -2,8 +2,8 @@ package discovery
2
3 import (
4 "errors"
5 + "fmt"
6 "reflect"
6 - "slices"
7 "sort"
8 "strings"
9 "sync"
@@ -17,28 +17,23 @@ import (
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
22 - relays map[string]RelayState
23 - policy RelayPolicy
24 - selfRelayKey string
25 - selfRelayURL string
20 + mu sync.RWMutex
21 + relays map[string]RelayState
22 + policy RelayPolicy
23 + self RelayState
24 }
25
26 func NewRelaySet(identity types.Identity, relayURL string, bootstrapRelayURLs []string) (*RelaySet, error) {
29 - if relayURL != "" {
30 - normalized, err := utils.NormalizeRelayURL(relayURL)
31 - if err != nil {
32 - return nil, err
33 - }
34 - relayURL = normalized
35 - }
36 -
27 set := &RelaySet{
38 - relays: make(map[string]RelayState),
39 - policy: DefaultRelayPolicy{},
40 - selfRelayKey: identity.Key(),
41 - selfRelayURL: relayURL,
28 + relays: make(map[string]RelayState),
29 + policy: DefaultRelayPolicy{},
30 + self: RelayState{
31 + Descriptor: types.RelayDescriptor{
32 + Identity: identity,
33 + RelayID: relayURL,
34 + APIHTTPSAddr: relayURL,
35 + },
36 + },
37 }
38 if err := set.SetBootstrapRelayURLs(bootstrapRelayURLs); err != nil {
39 return nil, err
@@ -56,72 +51,69 @@ func (s *RelaySet) SetRelayPolicy(policy RelayPolicy) {
51 }
52
53 func (s *RelaySet) SetBootstrapRelayURLs(inputs []string) error {
59 - normalized, err := utils.NormalizeRelayURLs(inputs...)
60 - if err != nil {
61 - return err
62 - }
63 -
54 s.mu.Lock()
55 defer s.mu.Unlock()
56
67 - filtered := utils.RemoveRelayURL(normalized, s.selfRelayURL)
68 -
57 + filtered := utils.RemoveRelayURL(inputs, s.self.Descriptor.APIHTTPSAddr)
58 keep := make(map[string]struct{}, len(filtered))
59 for _, relayURL := range filtered {
60 keep[relayURL] = struct{}{}
61 }
62
74 - for _, relayURL := range s.knownRelayURLs {
75 - if _, ok := keep[relayURL]; ok {
63 + seen := make(map[string]struct{}, len(filtered))
64 + for key, state := range s.relays {
65 + if state.Equal(s.self) {
66 + delete(s.relays, key)
67 continue
68 }
78 - state := s.relays[relayURL]
79 - if state.hasDescriptor() {
69 +
70 + _, bootstrap := keep[key]
71 + state.Bootstrap = bootstrap
72 + if !state.Bootstrap &&
73 + !state.hasDescriptor() &&
74 + !state.Banned &&
75 + state.consecutiveFailures == 0 {
76 + delete(s.relays, key)
77 continue
78 }
82 - if !s.policy.KeepState(state) {
83 - delete(s.relays, relayURL)
79 +
80 + s.relays[key] = state
81 + if bootstrap {
82 + seen[key] = struct{}{}
83 }
84 }
85
87 - s.knownRelayURLs = append([]string(nil), filtered...)
88 - return nil
89 -}
90 -
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
86 + for _, relayURL := range filtered {
87 + if _, ok := seen[relayURL]; ok {
88 + continue
89 }
90 +
91 + state := newRelayStateFromURL(relayURL)
92 + state.Bootstrap = true
93 + s.relays[relayURL] = state
94 }
100 - return ""
95 + return nil
96 }
97
103 -func (s *RelaySet) RelayStates() []RelayState {
98 +func (s *RelaySet) ActiveRelays() []RelayState {
99 s.mu.RLock()
100 defer s.mu.RUnlock()
101
107 - states := s.relayStatesLocked()
108 - for i := range states {
109 - states[i] = s.policy.Decide(states[i])
110 - }
111 - return states
102 + return s.policy.SelectActive(s.relayStatesLocked())
103 }
104
114 -func (s *RelaySet) ActiveRelays() []RelayState {
105 +func (s *RelaySet) OverlayPeerStates() []RelayState {
106 s.mu.RLock()
116 - defer s.mu.RUnlock()
117 -
107 states := s.relayStatesLocked()
108 + s.mu.RUnlock()
109 +
110 + now := time.Now().UTC()
111 out := make([]RelayState, 0, len(states))
112 for _, state := range states {
121 - state = s.policy.Decide(state)
122 - if state.Active && state.Descriptor.APIHTTPSAddr != "" {
123 - out = append(out, state)
113 + if !state.discoverable(now) || !state.Descriptor.SupportsOverlayPeer {
114 + continue
115 }
116 + out = append(out, state)
117 }
118 if len(out) == 0 {
119 return nil
@@ -129,120 +121,103 @@ func (s *RelaySet) ActiveRelays() []RelayState {
121 return out
122 }
123
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 -
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{}{}
124 +func (s *RelaySet) AdvertisedDescriptors() []types.RelayDescriptor {
125 + s.mu.RLock()
126 + defer s.mu.RUnlock()
127
145 - state, ok := s.relays[relayURL]
146 - if !ok {
147 - state = newRelayHintState(relayURL)
148 - }
149 - if selfRelay(state, s.selfRelayKey, s.selfRelayURL) {
150 - continue
151 - }
152 - state.Bootstrap = true
153 - out = append(out, state)
128 + states := s.policy.SelectAdvertised(s.relayStatesLocked())
129 + out := make([]types.RelayDescriptor, 0, len(states))
130 + for _, state := range states {
131 + out = append(out, state.Descriptor)
132 }
155 -
156 - descriptorStates := make([]RelayState, 0, len(s.relays))
157 - for relayURL, record := range s.relays {
158 - if !record.hasDescriptor() {
159 - continue
160 - }
161 - relayURL = strings.TrimSpace(record.Descriptor.APIHTTPSAddr)
162 - if relayURL == "" {
163 - continue
164 - }
165 - if selfRelay(record, s.selfRelayKey, s.selfRelayURL) {
166 - continue
167 - }
168 - if _, ok := seen[relayURL]; ok {
169 - continue
170 - }
171 -
172 - record.Bootstrap = slices.Contains(s.knownRelayURLs, relayURL)
173 - descriptorStates = append(descriptorStates, record)
133 + if len(out) == 0 {
134 + return nil
135 }
175 - sort.Slice(descriptorStates, func(i, j int) bool {
176 - return descriptorStates[i].Descriptor.APIHTTPSAddr < descriptorStates[j].Descriptor.APIHTTPSAddr
136 + sort.Slice(out, func(i, j int) bool {
137 + return out[i].APIHTTPSAddr < out[j].APIHTTPSAddr
138 })
178 - out = append(out, descriptorStates...)
139 + return out
140 +}
141
142 +func (s *RelaySet) relayStatesLocked() []RelayState {
143 + out := make([]RelayState, 0, len(s.relays))
144 + for _, state := range s.relays {
145 + out = append(out, state)
146 + }
147 if len(out) == 0 {
148 return nil
149 }
150 + sort.Slice(out, func(i, j int) bool {
151 + return out[i].Descriptor.APIHTTPSAddr < out[j].Descriptor.APIHTTPSAddr
152 + })
153 return out
154 }
155
186 -func (s *RelaySet) AdvertisedDescriptors() []types.RelayDescriptor {
187 - s.mu.RLock()
188 - defer s.mu.RUnlock()
189 -
190 - return s.policy.AdvertisedDescriptors(s.relayStatesLocked())
191 -}
192 -
156 func (s *RelaySet) BanRelayURL(relayURL string) {
157 s.mu.Lock()
158 defer s.mu.Unlock()
196 - relayURL = strings.TrimSpace(relayURL)
197 - if relayURL == "" {
198 - return
199 - }
159
201 - state := s.relays[relayURL]
202 - if state.Descriptor.APIHTTPSAddr == "" {
203 - state = newRelayHintState(relayURL)
160 + state, ok := s.relays[relayURL]
161 + if !ok {
162 + state = newRelayStateFromURL(relayURL)
163 + }
164 + if state.Equal(s.self) {
165 + delete(s.relays, relayURL)
166 + return
167 }
168 state = s.policy.OnBanned(state)
169 s.relays[relayURL] = state
170 }
171
209 -func (s *RelaySet) storeDescriptorLocked(state RelayState, policy RelayPolicy) (bool, bool, error) {
210 - desc := state.Descriptor
211 - relayKey := desc.Key()
212 - if relayKey == "" {
213 - return false, false, errors.New("descriptor identity is required")
214 - }
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")
172 +func (s *RelaySet) applyDiscoveredStateLocked(state RelayState, confirmed bool) (bool, error) {
173 + relayURL := state.Descriptor.APIHTTPSAddr
174 + relayKey := state.Descriptor.Key()
175 +
176 + previousState, hadPrevious := s.relays[relayURL]
177 + record := previousState
178 +
179 + if hadPrevious && record.hasDescriptor() && record.Descriptor.Key() != relayKey {
180 + return false, errors.New("descriptor identity does not match known relay url")
181 }
182
219 - previousURL := s.relayURLForKeyLocked(relayKey)
220 - record := s.relays[desc.APIHTTPSAddr]
221 - previousDescriptor := record.Descriptor
222 - if previousURL != "" && previousURL != desc.APIHTTPSAddr {
223 - previous := s.relays[previousURL]
224 - previousDescriptor = previous.Descriptor
225 - if !policy.KeepState(record) {
226 - record.Status = previous.Status
227 - record.consecutiveFailures = previous.consecutiveFailures
183 + previousURL := ""
184 + for url, existing := range s.relays {
185 + if url == relayURL || !existing.hasDescriptor() || existing.Descriptor.Key() != relayKey {
186 + continue
187 }
229 - delete(s.relays, previousURL)
188 + previousURL = url
189 + if !record.hasDescriptor() &&
190 + !record.Banned &&
191 + record.consecutiveFailures == 0 {
192 + record.Reachable = existing.Reachable
193 + record.Confirmed = existing.Confirmed
194 + record.Banned = existing.Banned
195 + record.consecutiveFailures = existing.consecutiveFailures
196 + }
197 + break
198 }
199
232 - added := previousURL == ""
233 - if !policy.KeepState(record) {
200 + if !record.hasDescriptor() &&
201 + !record.Banned &&
202 + record.consecutiveFailures == 0 {
203 record = state
204 } else {
236 - if record.FirstSeenAt.IsZero() {
237 - record.FirstSeenAt = state.FirstSeenAt
238 - }
239 - record.Descriptor = desc
205 + record.Descriptor = state.Descriptor
206 record.LastSeenAt = state.LastSeenAt
207 }
242 - s.relays[desc.APIHTTPSAddr] = record
208
244 - changed := added || !reflect.DeepEqual(previousDescriptor, desc)
245 - return added, changed, nil
209 + if confirmed {
210 + record = s.policy.OnConfirmed(record)
211 + } else {
212 + record = s.policy.OnHinted(record)
213 + }
214 +
215 + if previousURL != "" {
216 + delete(s.relays, previousURL)
217 + }
218 + s.relays[relayURL] = record
219 +
220 + return !hadPrevious || previousURL != "" || !reflect.DeepEqual(previousState, record), nil
221 }
222
223 func (s *RelaySet) ApplyRelayDiscoveryResponse(targetIdentity types.Identity, targetURL string, resp types.DiscoveryResponse, now time.Time) (relaySetChanged bool, err error) {
@@ -255,39 +230,63 @@ func (s *RelaySet) ApplyRelayDiscoveryResponse(targetIdentity types.Identity, ta
230 s.mu.Lock()
231 defer s.mu.Unlock()
232
258 - policy := s.policy
233 + if resp.ProtocolVersion != types.ProtocolVersion {
234 + return false, fmt.Errorf("relay protocol version mismatch: relay=%q client=%q", resp.ProtocolVersion, types.ProtocolVersion)
235 + }
236
260 - selfState, relayStates, err := policy.DiscoveryStates(targetIdentity, targetURL, s.selfRelayKey, s.selfRelayURL, resp, now)
237 + selfState, err := newRelayState(resp.Self, now)
238 if err != nil {
239 return false, err
240 }
264 -
265 - apply := func(state RelayState, advertise bool) error {
266 - desc := state.Descriptor
267 - added, descriptorChanged, err := s.storeDescriptorLocked(state, policy)
241 + if selfState.Equal(s.self) {
242 + return false, nil
243 + }
244 + if strings.TrimSpace(targetIdentity.Name) == "" && strings.TrimSpace(targetIdentity.Address) == "" {
245 + return false, errors.New("target relay identity is required")
246 + }
247 + if targetName := strings.TrimSpace(targetIdentity.Name); targetName != "" {
248 + if selfState.Descriptor.Name != utils.NormalizeHostname(targetName) {
249 + return false, errors.New("descriptor name does not match target relay")
250 + }
251 + }
252 + if targetAddress := strings.TrimSpace(targetIdentity.Address); targetAddress != "" {
253 + normalizedTargetAddress, err := utils.NormalizeEVMAddress(targetAddress)
254 if err != nil {
269 - return err
255 + return false, err
256 }
271 -
272 - storedState := s.relays[desc.APIHTTPSAddr]
273 - previousState := storedState
274 - storedState = policy.OnDiscovered(storedState, advertise)
275 - s.relays[desc.APIHTTPSAddr] = storedState
276 -
277 - changed := added || descriptorChanged || !reflect.DeepEqual(previousState, s.relays[desc.APIHTTPSAddr])
278 - if changed {
279 - relaySetChanged = true
257 + if selfState.Descriptor.Address != normalizedTargetAddress {
258 + return false, errors.New("descriptor address does not match target relay")
259 }
281 - return nil
260 }
283 -
284 - if err := apply(selfState, true); err != nil {
261 + if targetURL != "" && selfState.Descriptor.APIHTTPSAddr != strings.TrimSpace(targetURL) {
262 + return false, errors.New("descriptor api_https_addr does not match target url")
263 + }
264 + changed, err := s.applyDiscoveredStateLocked(selfState, true)
265 + if err != nil {
266 return false, err
267 }
287 - for _, relayState := range relayStates {
288 - if err := apply(relayState, false); err != nil {
268 + relaySetChanged = relaySetChanged || changed
269 +
270 + seen := map[string]struct{}{selfState.Descriptor.Key(): {}}
271 + for _, descriptor := range resp.Relays {
272 + relayState, err := newRelayState(descriptor, now)
273 + if err != nil {
274 + continue
275 + }
276 + if relayState.Equal(s.self) {
277 + continue
278 + }
279 + relayKey := relayState.Descriptor.Key()
280 + if _, ok := seen[relayKey]; ok {
281 + continue
282 + }
283 + seen[relayKey] = struct{}{}
284 +
285 + changed, err := s.applyDiscoveredStateLocked(relayState, false)
286 + if err != nil {
287 return false, err
288 }
289 + relaySetChanged = relaySetChanged || changed
290 }
291 return relaySetChanged, nil
292 }
@@ -297,22 +296,29 @@ func (s *RelaySet) RecordDiscoveryFailure(identity types.Identity, relayURL stri
296 if relayKey == "" {
297 return false, "", 0
298 }
300 - relayURL = strings.TrimSpace(relayURL)
299
300 s.mu.Lock()
301 defer s.mu.Unlock()
302
305 - if relayURL == "" || s.relays[relayURL].Descriptor.Key() != relayKey {
306 - relayURL = s.relayURLForKeyLocked(relayKey)
303 + state, ok := s.relays[relayURL]
304 + if ok && state.hasDescriptor() && state.Descriptor.Key() != relayKey {
305 + ok = false
306 }
308 - if relayURL == "" {
309 - return false, "", 0
307 + if !ok {
308 + for url, existing := range s.relays {
309 + if !existing.hasDescriptor() || existing.Descriptor.Key() != relayKey {
310 + continue
311 + }
312 + relayURL = url
313 + state = existing
314 + ok = true
315 + break
316 + }
317 }
311 -
312 - state, ok := s.relays[relayURL]
318 if !ok {
319 return false, "", 0
320 }
321 +
322 state, expired, expireReason = s.policy.OnFailure(state, err, recoveryFailures)
323 s.relays[relayURL] = state
324 return expired, expireReason, state.consecutiveFailures
portal/discovery/relaystate.go new
+93
@@ -0,0 +1,93 @@
1 +package discovery
2 +
3 +import (
4 + "errors"
5 + "strings"
6 + "time"
7 +
8 + "github.com/gosuda/portal-tunnel/v2/types"
9 + "github.com/gosuda/portal-tunnel/v2/utils"
10 +)
11 +
12 +type RelayState struct {
13 + Descriptor types.RelayDescriptor
14 + Bootstrap bool
15 + Reachable bool
16 + Confirmed bool
17 + Banned bool
18 + LastSeenAt time.Time
19 +
20 + consecutiveFailures int
21 +}
22 +
23 +func newRelayState(desc types.RelayDescriptor, seenAt time.Time) (RelayState, error) {
24 + state := RelayState{
25 + Descriptor: desc,
26 + }
27 + if seenAt.IsZero() {
28 + return state, nil
29 + }
30 +
31 + seenAt = seenAt.UTC()
32 + normalized, err := utils.NormalizeDescriptor(desc)
33 + if err != nil {
34 + return RelayState{}, err
35 + }
36 + if normalized.ExpiresAt.Before(seenAt) {
37 + return RelayState{}, errors.New("descriptor expired")
38 + }
39 +
40 + state.Descriptor = normalized
41 + state.LastSeenAt = seenAt
42 + return state, nil
43 +}
44 +
45 +func newRelayStateFromURL(relayURL string) RelayState {
46 + return RelayState{
47 + Descriptor: types.RelayDescriptor{
48 + Identity: types.Identity{
49 + Name: utils.PortalRootHost(relayURL),
50 + },
51 + RelayID: relayURL,
52 + APIHTTPSAddr: relayURL,
53 + },
54 + }
55 +}
56 +
57 +func (state RelayState) hasDescriptor() bool {
58 + return !state.LastSeenAt.IsZero()
59 +}
60 +
61 +func (state RelayState) discoverable(now time.Time) bool {
62 + if state.Banned {
63 + return false
64 + }
65 + if !state.hasDescriptor() {
66 + return state.Bootstrap
67 + }
68 + if !state.Bootstrap && !state.Reachable {
69 + return false
70 + }
71 + if !state.Descriptor.ExpiresAt.After(now) {
72 + return false
73 + }
74 + if state.Descriptor.SupportsOverlayPeer &&
75 + (state.Descriptor.WireGuardPublicKey == "" ||
76 + state.Descriptor.WireGuardEndpoint == "" ||
77 + state.Descriptor.OverlayIPv4 == "") {
78 + return false
79 + }
80 + return true
81 +}
82 +
83 +func (state RelayState) Equal(other RelayState) bool {
84 + stateKey := state.Descriptor.Key()
85 + otherKey := other.Descriptor.Key()
86 + if stateKey != "" && otherKey != "" && stateKey == otherKey {
87 + return true
88 + }
89 +
90 + stateURL := strings.TrimSpace(state.Descriptor.APIHTTPSAddr)
91 + otherURL := strings.TrimSpace(other.Descriptor.APIHTTPSAddr)
92 + return stateURL != "" && otherURL != "" && stateURL == otherURL
93 +}
portal/overlay/overlay.go
+1 -1
@@ -219,7 +219,7 @@ func (o *Overlay) Sync(relays []discovery.RelayState) error {
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 {
222 + if !relay.Descriptor.SupportsOverlayPeer {
223 continue
224 }
225
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.RelayStates()); err != nil {
617 + if err := overlay.Sync(s.relaySet.OverlayPeerStates()); err != nil {
618 _ = overlay.Shutdown(context.Background())
619 return nil, fmt.Errorf("sync wireguard peers: %w", err)
620 }
portal/server_test.go
+6 -12
@@ -509,10 +509,8 @@ func TestServerSetBootstrapRelayURLsAllowsLoopbackButSkipsSelfRelay(t *testing.T
509 }
510 advertisedDescriptors := server.relaySet.AdvertisedDescriptors()
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 - }
512 + for _, state := range server.relaySet.ActiveRelays() {
513 + knownURLs = append(knownURLs, state.Descriptor.APIHTTPSAddr)
514 }
515 sort.Strings(knownURLs)
516 if !reflect.DeepEqual(knownURLs, []string{
@@ -567,10 +565,8 @@ func TestServerDiscoverySkipsSelfRelayHint(t *testing.T) {
565 }
566
567 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 - }
568 + for _, state := range server.relaySet.ActiveRelays() {
569 + knownURLs = append(knownURLs, state.Descriptor.APIHTTPSAddr)
570 }
571 sort.Strings(knownURLs)
572 if !reflect.DeepEqual(knownURLs, []string{"https://bootstrap.example.com"}) {
@@ -617,10 +613,8 @@ func TestServerRecordVerifiedDiscoveryPeerRequiresDirectConfirmation(t *testing.
613 t.Fatalf("applyRelayDiscoveryResponse() hinted error = %v", err)
614 }
615 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 - }
616 + for _, state := range server.relaySet.ActiveRelays() {
617 + knownURLs = append(knownURLs, state.Descriptor.APIHTTPSAddr)
618 }
619 sort.Strings(knownURLs)
620 if !reflect.DeepEqual(knownURLs, []string{"https://bootstrap.example.com"}) {
sdk/mitm_test.go
+4 -6
@@ -228,8 +228,8 @@ func TestMITMProbeDetectionBansListener(t *testing.T) {
228 Reason: types.MITMProbeReasonExporterMismatch,
229 }, nil)
230
231 - for _, state := range listener.relaySet.RelayStates() {
232 - if state.Active && state.Descriptor.APIHTTPSAddr == relayURL.String() {
231 + for _, state := range listener.relaySet.ActiveRelays() {
232 + if state.Descriptor.APIHTTPSAddr == relayURL.String() {
233 t.Fatal("relay still active after mitm detection")
234 }
235 }
@@ -261,10 +261,8 @@ func TestMITMProbeDetectionWarnsWithoutBanningListener(t *testing.T) {
261 }, nil)
262
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 - }
264 + for _, state := range listener.relaySet.ActiveRelays() {
265 + activeRelayURLs = append(activeRelayURLs, state.Descriptor.APIHTTPSAddr)
266 }
267 if len(activeRelayURLs) != 1 || activeRelayURLs[0] != relayURL.String() {
268 t.Fatalf("ActiveRelayURLs() = %v, want [%q]", activeRelayURLs, relayURL.String())
utils/identity.go
+14
@@ -86,6 +86,20 @@ func NormalizeDescriptor(desc types.RelayDescriptor) (types.RelayDescriptor, err
86 if desc.SignerPublicKey == "" {
87 desc.SignerPublicKey = desc.PublicKey
88 }
89 +
90 + switch {
91 + case desc.Name == "":
92 + return types.RelayDescriptor{}, errors.New("identity.name is required")
93 + case desc.APIHTTPSAddr == "":
94 + return types.RelayDescriptor{}, errors.New("api_https_addr is required")
95 + case desc.RelayID != desc.APIHTTPSAddr:
96 + return types.RelayDescriptor{}, errors.New("relay_id must match api_https_addr")
97 + case desc.ExpiresAt.IsZero():
98 + return types.RelayDescriptor{}, errors.New("expires_at is required")
99 + case desc.IssuedAt.After(desc.ExpiresAt):
100 + return types.RelayDescriptor{}, errors.New("issued_at must be before expires_at")
101 + }
102 +
103 return desc, nil
104 }
105