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