feat: enhance relay discovery and selection logic, add support for discovery flag in relay descriptors

Kim committed Apr 13, 2026 at 17:04 UTC 7a0608be29c6b095f86a5dbfe41d058f12b55a72
12 files changed +253 -323
portal/api_server.go
+2 -15
@@ -152,10 +152,6 @@ func (s *Server) extractAllowedClientIP(w http.ResponseWriter, r *http.Request)
152 }
153
154 func (s *Server) handleRelayDiscovery(w http.ResponseWriter, r *http.Request) {
155 - if !utils.RequireMethod(w, r, http.MethodGet) {
156 - return
157 - }
158 -
155 now := time.Now().UTC()
156 activeConns := float64(s.proxy.ActiveConns())
157 tcpTrafficBPS := s.proxy.CurrentTCPBPS(now)
@@ -182,6 +178,7 @@ func (s *Server) handleRelayDiscovery(w http.ResponseWriter, r *http.Request) {
178 IssuedAt: now,
179 ExpiresAt: now.Add(2 * discovery.DiscoveryPollInterval),
180 APIHTTPSAddr: s.cfg.PortalURL,
181 + Discovery: s.cfg.DiscoveryEnabled,
182 IngressTLSAddr: ingressAddr,
183 WireGuardPublicKey: wireGuardPublicKey,
184 WireGuardEndpoint: wireGuardEndpoint,
@@ -198,17 +195,7 @@ func (s *Server) handleRelayDiscovery(w http.ResponseWriter, r *http.Request) {
195 utils.WriteAPIError(w, http.StatusInternalServerError, types.APIErrorCodeInternal, err.Error())
196 return
197 }
201 -
202 - resp := types.DiscoveryResponse{
203 - ProtocolVersion: types.ProtocolVersion,
204 - GeneratedAt: now,
205 - Self: self,
206 - Relays: nil,
207 - }
208 - if s.relaySet != nil {
209 - resp.Relays = s.relaySet.ConfirmedDescriptors()
210 - }
211 - utils.WriteAPIData(w, http.StatusOK, resp)
198 + s.relaySet.ServeDiscovery(w, r, self)
199 }
200
201 func (s *Server) handleDomain(w http.ResponseWriter, r *http.Request) {
portal/discovery/policy.go
+2 -10
@@ -5,7 +5,6 @@ import (
5 "math/rand"
6 "net/http"
7 "slices"
8 - "strings"
8 "time"
9
10 "github.com/gosuda/portal-tunnel/v2/types"
@@ -13,7 +12,6 @@ import (
12
13 type RelayPolicy interface {
14 SelectActive([]RelayState) []RelayState
16 - SelectConfirmed([]RelayState) []RelayState
15 SelectPriority([]RelayState, ClientState) []string
16 OnConfirmed(RelayState) RelayState
17 OnHinted(RelayState) RelayState
@@ -49,12 +47,6 @@ func (p DefaultRelayPolicy) SelectActive(states []RelayState) []RelayState {
47 })
48 }
49
52 -func (p DefaultRelayPolicy) SelectConfirmed(states []RelayState) []RelayState {
53 - return p.selectStates(states, func(state RelayState) bool {
54 - return state.Confirmed
55 - })
56 -}
57 -
50 func (p DefaultRelayPolicy) SelectPriority(states []RelayState, clientState ClientState) []string {
51 selected := p.SelectActive(states)
52 if len(selected) == 0 {
@@ -70,7 +62,7 @@ func (p DefaultRelayPolicy) SelectPriority(states []RelayState, clientState Clie
62 if clientState.RequireTCP && state.hasDescriptor() && !state.Descriptor.SupportsTCP {
63 continue
64 }
73 - relayURL := strings.TrimSpace(state.Descriptor.APIHTTPSAddr)
65 + relayURL := state.Descriptor.APIHTTPSAddr
66 if slices.Contains(clientState.ExplicitRelayURLs, relayURL) {
67 explicit = append(explicit, relayURL)
68 continue
@@ -86,7 +78,7 @@ func (p DefaultRelayPolicy) SelectPriority(states []RelayState, clientState Clie
78 highRTTAuto := make([]string, 0, len(autoPool))
79 penalizedAuto := make([]string, 0, len(autoPool))
80 for _, state := range autoPool {
89 - relayURL := strings.TrimSpace(state.Descriptor.APIHTTPSAddr)
81 + relayURL := state.Descriptor.APIHTTPSAddr
82 statePenalized := state.consecutiveFailures > 0 && !state.Reachable
83 switch {
84 case slices.Contains(clientState.ActiveRelayURLs, relayURL) && !statePenalized:
portal/discovery/refresher.go
+37 -43
@@ -62,7 +62,7 @@ func NewRefresher(relaySet *RelaySet, rootCAPEM []byte, overlay OverlayRuntime)
62 }, nil
63 }
64
65 -func (r *Refresher) Refresh(ctx context.Context) error {
65 +func (r *Refresher) Refresh(ctx context.Context, extraSourceURLs ...string) error {
66 if r.overlay != nil {
67 if err := r.refreshOverlay(ctx); err != nil && ctx.Err() == nil {
68 log.Warn().
@@ -73,27 +73,40 @@ func (r *Refresher) Refresh(ctx context.Context) error {
73 return ctx.Err()
74 }
75 }
76 - return r.refreshHTTPS(ctx)
76 + return r.refreshHTTPS(ctx, extraSourceURLs)
77 }
78
79 -func (r *Refresher) refreshHTTPS(ctx context.Context) error {
79 +func (r *Refresher) refreshHTTPS(ctx context.Context, extraSourceURLs []string) error {
80 r.relaySet.mu.RLock()
81 states := r.relaySet.relayStatesLocked()
82 r.relaySet.mu.RUnlock()
83
84 now := time.Now().UTC()
85 for _, state := range states {
86 - if !state.discoverable(now) || !state.Bootstrap {
86 + if !state.discoverable(now) || (state.hasDescriptor() && !state.Descriptor.Discovery) {
87 continue
88 }
89 - relay := state.Descriptor
90 - baseURL, err := url.Parse(relay.APIHTTPSAddr)
89 +
90 + relayURL := state.Descriptor.APIHTTPSAddr
91 + if relayURL == "" {
92 + continue
93 + }
94 +
95 + recoveryFailures := r.directRecoveryFailures
96 + if state.Bootstrap {
97 + recoveryFailures = 0
98 + }
99 +
100 + baseURL, err := url.Parse(relayURL)
101 if err != nil {
102 + if recoveryFailures > 0 {
103 + r.logDiscoveryFailure(relayURL, relayURL, recoveryFailures, err)
104 + }
105 continue
106 }
107 if utils.IsLocalRelayHost(baseURL.Hostname()) {
108 log.Info().
96 - Str("relay", relay.APIHTTPSAddr).
109 + Str("relay", relayURL).
110 Msg("skip loopback relay as discovery source")
111 continue
112 }
@@ -104,57 +117,38 @@ func (r *Refresher) refreshHTTPS(ctx context.Context) error {
117 if ctx.Err() != nil {
118 return ctx.Err()
119 }
120 + if recoveryFailures > 0 {
121 + r.logDiscoveryFailure(relayURL, relayURL, recoveryFailures, err)
122 + }
123 continue
124 }
109 -
125 measuredAt := time.Now().UTC()
111 - if _, err := r.relaySet.ApplyRelayDiscoveryResponse(relay.Identity, relay.APIHTTPSAddr, resp, measuredAt); err != nil {
126 +
127 + if _, err := r.relaySet.ApplyRelayDiscoveryResponse(relayURL, resp, measuredAt); err != nil {
128 + if recoveryFailures > 0 {
129 + r.logDiscoveryFailure(relayURL, relayURL, recoveryFailures, err)
130 + }
131 continue
132 }
114 - r.relaySet.RecordDiscoveryRTT(relay.APIHTTPSAddr, time.Since(startedAt), measuredAt)
115 - }
116 - if err := ctx.Err(); err != nil {
117 - return err
133 + r.relaySet.RecordDiscoveryRTT(relayURL, time.Since(startedAt), measuredAt)
134 }
135
120 - r.relaySet.mu.RLock()
121 - states = r.relaySet.relayStatesLocked()
122 - r.relaySet.mu.RUnlock()
123 -
124 - now = time.Now().UTC()
125 - for _, state := range states {
126 - if !state.discoverable(now) || state.Bootstrap {
127 - continue
128 - }
129 - relay := state.Descriptor
130 - baseURL, err := url.Parse(relay.APIHTTPSAddr)
136 + for _, sourceURL := range extraSourceURLs {
137 + baseURL, err := url.Parse(sourceURL)
138 if err != nil {
132 - r.logDirectDiscoveryFailure(relay, err, r.directRecoveryFailures)
133 - continue
134 - }
135 - if utils.IsLocalRelayHost(baseURL.Hostname()) {
136 - log.Info().
137 - Str("relay", relay.APIHTTPSAddr).
138 - Msg("skip loopback relay as discovery source")
139 continue
140 }
141
142 - startedAt := time.Now()
142 var resp types.DiscoveryResponse
143 if err := utils.HTTPDoAPIPath(ctx, r.httpClient, baseURL, http.MethodGet, types.PathDiscovery, nil, nil, &resp); err != nil {
144 if ctx.Err() != nil {
145 return ctx.Err()
146 }
148 - r.logDirectDiscoveryFailure(relay, err, r.directRecoveryFailures)
147 continue
148 }
151 -
152 - measuredAt := time.Now().UTC()
153 - if _, err := r.relaySet.ApplyRelayDiscoveryResponse(relay.Identity, relay.APIHTTPSAddr, resp, measuredAt); err != nil {
154 - r.logDirectDiscoveryFailure(relay, err, r.directRecoveryFailures)
149 + if _, err := r.relaySet.ApplyRelayDiscoveryResponse("", resp, time.Now().UTC()); err != nil {
150 continue
151 }
157 - r.relaySet.RecordDiscoveryRTT(relay.APIHTTPSAddr, time.Since(startedAt), measuredAt)
152 }
153 return nil
154 }
@@ -174,7 +168,7 @@ func (r *Refresher) refreshOverlay(ctx context.Context) error {
168 return err
169 }
170
177 - relaySetChanged, err := r.relaySet.ApplyRelayDiscoveryResponse(relay.Identity, relay.APIHTTPSAddr, resp, time.Now().UTC())
171 + relaySetChanged, err := r.relaySet.ApplyRelayDiscoveryResponse(relay.APIHTTPSAddr, resp, time.Now().UTC())
172 if err != nil {
173 return err
174 }
@@ -188,19 +182,19 @@ func (r *Refresher) refreshOverlay(ctx context.Context) error {
182 return nil
183 }
184
191 -func (r *Refresher) logDirectDiscoveryFailure(relay types.RelayDescriptor, err error, recoveryFailures int) {
192 - expired, expireReason, consecutiveFailures := r.relaySet.RecordDiscoveryFailure(relay.Identity, relay.APIHTTPSAddr, err, recoveryFailures)
185 +func (r *Refresher) logDiscoveryFailure(targetRelayURL, sourceURL string, recoveryFailures int, err error) {
186 + expired, expireReason, consecutiveFailures := r.relaySet.RecordRelayFailure(targetRelayURL, err, recoveryFailures)
187 if !expired {
188 return
189 }
190
191 event := log.Warn().
192 Err(err).
199 - Str("relay", relay.APIHTTPSAddr).
193 + Str("relay", sourceURL).
194 Bool("expired", true).
195 Str("reason", expireReason)
196 if consecutiveFailures > 0 {
197 event = event.Int("consecutive_failures", consecutiveFailures)
198 }
205 - event.Msg("direct relay discovery expired")
199 + event.Msg("discovery source expired")
200 }
portal/discovery/relayset.go
+92 -142
@@ -3,9 +3,9 @@ package discovery
3 import (
4 "errors"
5 "fmt"
6 + "net/http"
7 "reflect"
8 "sort"
8 - "strings"
9 "sync"
10 "time"
11
@@ -54,10 +54,7 @@ func (s *RelaySet) SetBootstrapRelayURLs(inputs []string) error {
54 for key, state := range s.relays {
55 _, bootstrap := keep[key]
56 state.Bootstrap = bootstrap
57 - if !state.Bootstrap &&
58 - !state.hasDescriptor() &&
59 - !state.Banned &&
60 - state.consecutiveFailures == 0 {
57 + if !state.Bootstrap && !state.hasDescriptor() && !state.Banned && state.consecutiveFailures == 0 {
58 delete(s.relays, key)
59 continue
60 }
@@ -115,13 +112,17 @@ func (s *RelaySet) OverlayPeerStates() []RelayState {
112 return out
113 }
114
118 -func (s *RelaySet) ConfirmedDescriptors() []types.RelayDescriptor {
115 +func (s *RelaySet) Descriptors() []types.RelayDescriptor {
116 s.mu.RLock()
120 - defer s.mu.RUnlock()
117 + states := s.relayStatesLocked()
118 + s.mu.RUnlock()
119
122 - states := s.policy.SelectConfirmed(s.relayStatesLocked())
120 + now := time.Now().UTC()
121 out := make([]types.RelayDescriptor, 0, len(states))
122 for _, state := range states {
123 + if !state.hasDescriptor() || !state.Descriptor.ExpiresAt.After(now) || !state.Descriptor.Discovery {
124 + continue
125 + }
126 out = append(out, state.Descriptor)
127 }
128 if len(out) == 0 {
@@ -130,6 +131,40 @@ func (s *RelaySet) ConfirmedDescriptors() []types.RelayDescriptor {
131 return out
132 }
133
134 +func (s *RelaySet) ServeDiscovery(w http.ResponseWriter, r *http.Request, local ...types.RelayDescriptor) {
135 + if !utils.RequireMethod(w, r, http.MethodGet) {
136 + return
137 + }
138 +
139 + known := s.Descriptors()
140 + relays := make([]types.RelayDescriptor, 0, len(local)+len(known))
141 + seen := make(map[string]struct{}, len(local)+len(known))
142 + add := func(descriptor types.RelayDescriptor) {
143 + relayURL := descriptor.APIHTTPSAddr
144 + if relayURL == "" {
145 + return
146 + }
147 + if _, ok := seen[relayURL]; ok {
148 + return
149 + }
150 + seen[relayURL] = struct{}{}
151 + relays = append(relays, descriptor)
152 + }
153 +
154 + for _, descriptor := range local {
155 + add(descriptor)
156 + }
157 + for _, descriptor := range known {
158 + add(descriptor)
159 + }
160 +
161 + utils.WriteAPIData(w, http.StatusOK, types.DiscoveryResponse{
162 + ProtocolVersion: types.ProtocolVersion,
163 + GeneratedAt: time.Now().UTC(),
164 + Relays: relays,
165 + })
166 +}
167 +
168 func (s *RelaySet) relayStatesLocked() []RelayState {
169 out := make([]RelayState, 0, len(s.relays))
170 for _, state := range s.relays {
@@ -156,123 +191,72 @@ func (s *RelaySet) BanRelayURL(relayURL string) {
191 s.relays[relayURL] = state
192 }
193
159 -func (s *RelaySet) applyDiscoveredStateLocked(state RelayState, confirmed bool) (bool, error) {
160 - relayURL := state.Descriptor.APIHTTPSAddr
161 - relayKey := state.Descriptor.Key()
162 -
163 - previousState, hadPrevious := s.relays[relayURL]
164 - record := previousState
165 - bootstrap := record.Bootstrap
166 -
167 - if hadPrevious && record.hasDescriptor() && record.Descriptor.Key() != relayKey {
168 - return false, errors.New("descriptor identity does not match known relay url")
169 - }
170 -
171 - previousURL := ""
172 - for url, existing := range s.relays {
173 - if url == relayURL || !existing.hasDescriptor() || existing.Descriptor.Key() != relayKey {
174 - continue
175 - }
176 - previousURL = url
177 - bootstrap = bootstrap || existing.Bootstrap
178 - if !record.hasDescriptor() &&
179 - !record.Banned &&
180 - record.consecutiveFailures == 0 {
181 - record.Reachable = existing.Reachable
182 - record.Confirmed = existing.Confirmed
183 - record.Banned = existing.Banned
184 - record.DiscoveryRTT = existing.DiscoveryRTT
185 - record.DiscoveryRTTAt = existing.DiscoveryRTTAt
186 - record.consecutiveFailures = existing.consecutiveFailures
187 - }
188 - break
189 - }
190 -
191 - if !record.hasDescriptor() &&
192 - !record.Banned &&
193 - record.consecutiveFailures == 0 {
194 - record = state
195 - } else {
196 - record.Descriptor = state.Descriptor
197 - record.LastSeenAt = state.LastSeenAt
198 - }
199 - record.Bootstrap = bootstrap
200 -
201 - if confirmed {
202 - record = s.policy.OnConfirmed(record)
203 - } else {
204 - record = s.policy.OnHinted(record)
205 - }
206 -
207 - if previousURL != "" {
208 - delete(s.relays, previousURL)
209 - }
210 - s.relays[relayURL] = record
211 -
212 - return !hadPrevious || previousURL != "" || !reflect.DeepEqual(previousState, record), nil
213 -}
214 -
215 -func (s *RelaySet) ApplyRelayDiscoveryResponse(targetIdentity types.Identity, targetURL string, resp types.DiscoveryResponse, now time.Time) (relaySetChanged bool, err error) {
194 +func (s *RelaySet) ApplyRelayDiscoveryResponse(targetURL string, resp types.DiscoveryResponse, now time.Time) (relaySetChanged bool, err error) {
195 if now.IsZero() {
196 now = time.Now().UTC()
197 } else {
198 now = now.UTC()
199 }
200
222 - s.mu.Lock()
223 - defer s.mu.Unlock()
224 -
201 if resp.ProtocolVersion != types.ProtocolVersion {
202 return false, fmt.Errorf("relay protocol version mismatch: relay=%q client=%q", resp.ProtocolVersion, types.ProtocolVersion)
203 }
204 + authoritative := targetURL != ""
205
229 - selfState, err := newRelayState(resp.Self, now)
230 - if err != nil {
231 - return false, err
232 - }
233 - if strings.TrimSpace(targetIdentity.Name) == "" && strings.TrimSpace(targetIdentity.Address) == "" {
234 - return false, errors.New("target relay identity is required")
235 - }
236 - if targetName := strings.TrimSpace(targetIdentity.Name); targetName != "" {
237 - if selfState.Descriptor.Name != utils.NormalizeHostname(targetName) {
238 - return false, errors.New("descriptor name does not match target relay")
239 - }
240 - }
241 - if targetAddress := strings.TrimSpace(targetIdentity.Address); targetAddress != "" {
242 - normalizedTargetAddress, err := utils.NormalizeEVMAddress(targetAddress)
243 - if err != nil {
244 - return false, err
245 - }
246 - if selfState.Descriptor.Address != normalizedTargetAddress {
247 - return false, errors.New("descriptor address does not match target relay")
248 - }
249 - }
250 - if targetURL != "" && selfState.Descriptor.APIHTTPSAddr != strings.TrimSpace(targetURL) {
251 - return false, errors.New("descriptor api_https_addr does not match target url")
252 - }
253 - changed, err := s.applyDiscoveredStateLocked(selfState, true)
254 - if err != nil {
255 - return false, err
256 - }
257 - relaySetChanged = relaySetChanged || changed
206 + s.mu.Lock()
207 + defer s.mu.Unlock()
208
259 - seen := map[string]struct{}{selfState.Descriptor.Key(): {}}
209 + discoveredByURL := make(map[string]RelayState, len(resp.Relays))
210 + discoveredOrder := make([]string, 0, len(resp.Relays))
211 + targetFound := false
212 for _, descriptor := range resp.Relays {
213 relayState, err := newRelayState(descriptor, now)
214 if err != nil {
215 continue
216 }
265 - relayKey := relayState.Descriptor.Key()
266 - if _, ok := seen[relayKey]; ok {
217 + relayURL := relayState.Descriptor.APIHTTPSAddr
218 + if relayURL == "" {
219 continue
220 }
269 - seen[relayKey] = struct{}{}
221 + if authoritative && relayURL == targetURL {
222 + targetFound = true
223 + }
224 + if _, ok := discoveredByURL[relayURL]; !ok {
225 + discoveredOrder = append(discoveredOrder, relayURL)
226 + }
227 + discoveredByURL[relayURL] = relayState
228 + }
229 +
230 + if authoritative && !targetFound {
231 + return false, errors.New("target relay descriptor missing from relays")
232 + }
233 +
234 + for _, relayURL := range discoveredOrder {
235 + record := discoveredByURL[relayURL]
236 + existingAtURL, hasExistingAtURL := s.relays[relayURL]
237 + record.Bootstrap = record.Bootstrap || existingAtURL.Bootstrap
238 + record.Reachable = record.Reachable || existingAtURL.Reachable
239 + record.Confirmed = record.Confirmed || existingAtURL.Confirmed
240 + record.Banned = record.Banned || existingAtURL.Banned
241 + if record.consecutiveFailures < existingAtURL.consecutiveFailures {
242 + record.consecutiveFailures = existingAtURL.consecutiveFailures
243 + }
244 + if record.DiscoveryRTTAt.IsZero() || (!existingAtURL.DiscoveryRTTAt.IsZero() && existingAtURL.DiscoveryRTTAt.After(record.DiscoveryRTTAt)) {
245 + record.DiscoveryRTT = existingAtURL.DiscoveryRTT
246 + record.DiscoveryRTTAt = existingAtURL.DiscoveryRTTAt
247 + }
248
271 - changed, err := s.applyDiscoveredStateLocked(relayState, false)
272 - if err != nil {
273 - return false, err
249 + if authoritative && relayURL == targetURL {
250 + record = s.policy.OnConfirmed(record)
251 + } else {
252 + record = s.policy.OnHinted(record)
253 + }
254 +
255 + s.relays[relayURL] = record
256 +
257 + if !hasExistingAtURL || !reflect.DeepEqual(existingAtURL, record) {
258 + relaySetChanged = true
259 }
275 - relaySetChanged = relaySetChanged || changed
260 }
261 return relaySetChanged, nil
262 }
@@ -291,12 +275,6 @@ func (s *RelaySet) RecordDiscoveryRTT(relayURL string, rtt time.Duration, measur
275 s.relays[relayURL] = state
276 }
277
294 -func (s *RelaySet) recordRelayFailureLocked(relayURL string, state RelayState, err error, recoveryFailures int) (expired bool, expireReason string, consecutiveFailures int) {
295 - state, expired, expireReason = s.policy.OnFailure(state, err, recoveryFailures)
296 - s.relays[relayURL] = state
297 - return expired, expireReason, state.consecutiveFailures
298 -}
299 -
278 func (s *RelaySet) RecordRelayFailure(relayURL string, err error, recoveryFailures int) (expired bool, expireReason string, consecutiveFailures int) {
279 s.mu.Lock()
280 defer s.mu.Unlock()
@@ -305,35 +283,7 @@ func (s *RelaySet) RecordRelayFailure(relayURL string, err error, recoveryFailur
283 if !ok {
284 return false, "", 0
285 }
308 - return s.recordRelayFailureLocked(relayURL, state, err, recoveryFailures)
309 -}
310 -
311 -func (s *RelaySet) RecordDiscoveryFailure(identity types.Identity, relayURL string, err error, recoveryFailures int) (expired bool, expireReason string, consecutiveFailures int) {
312 - relayKey := identity.Key()
313 - if relayKey == "" {
314 - return false, "", 0
315 - }
316 -
317 - s.mu.Lock()
318 - defer s.mu.Unlock()
319 -
320 - state, ok := s.relays[relayURL]
321 - if ok && state.hasDescriptor() && state.Descriptor.Key() != relayKey {
322 - ok = false
323 - }
324 - if !ok {
325 - for url, existing := range s.relays {
326 - if !existing.hasDescriptor() || existing.Descriptor.Key() != relayKey {
327 - continue
328 - }
329 - relayURL = url
330 - state = existing
331 - ok = true
332 - break
333 - }
334 - }
335 - if !ok {
336 - return false, "", 0
337 - }
338 - return s.recordRelayFailureLocked(relayURL, state, err, recoveryFailures)
286 + state, expired, expireReason = s.policy.OnFailure(state, err, recoveryFailures)
287 + s.relays[relayURL] = state
288 + return expired, expireReason, state.consecutiveFailures
289 }
portal/discovery/relayset_test.go
+2 -25
@@ -14,9 +14,9 @@ func TestApplyRelayDiscoveryResponsePreservesBootstrapFlag(t *testing.T) {
14 }
15
16 desc := mustPolicyRelayDescriptor(t, "relay-a", "https://relay-a.example")
17 - if _, err := set.ApplyRelayDiscoveryResponse(desc.Identity, desc.APIHTTPSAddr, types.DiscoveryResponse{
17 + if _, err := set.ApplyRelayDiscoveryResponse(desc.APIHTTPSAddr, types.DiscoveryResponse{
18 ProtocolVersion: types.ProtocolVersion,
19 - Self: desc,
19 + Relays: []types.RelayDescriptor{desc},
20 }, time.Now().UTC()); err != nil {
21 t.Fatalf("ApplyRelayDiscoveryResponse() error = %v", err)
22 }
@@ -29,26 +29,3 @@ func TestApplyRelayDiscoveryResponsePreservesBootstrapFlag(t *testing.T) {
29 t.Fatal("bootstrap relay lost bootstrap flag after discovery update")
30 }
31 }
32 -
33 -func TestApplyRelayDiscoveryResponseAllowsURLChangeForSameIdentity(t *testing.T) {
34 - set, err := NewRelaySet(nil)
35 - if err != nil {
36 - t.Fatalf("NewRelaySet() error = %v", err)
37 - }
38 -
39 - desc := mustPolicyRelayDescriptor(t, "relay-a", "https://relay-a.example")
40 - if _, err := set.ApplyRelayDiscoveryResponse(desc.Identity, desc.APIHTTPSAddr, types.DiscoveryResponse{
41 - ProtocolVersion: types.ProtocolVersion,
42 - Self: desc,
43 - }, time.Now().UTC()); err != nil {
44 - t.Fatalf("ApplyRelayDiscoveryResponse() error = %v", err)
45 - }
46 -
47 - changedURL := mustPolicyRelayDescriptor(t, desc.Name, "https://relay-b.example")
48 - if _, err := set.ApplyRelayDiscoveryResponse(desc.Identity, "", types.DiscoveryResponse{
49 - ProtocolVersion: types.ProtocolVersion,
50 - Self: changedURL,
51 - }, time.Now().UTC()); err != nil {
52 - t.Fatalf("ApplyRelayDiscoveryResponse() error = %v, want nil for same relay identity", err)
53 - }
54 -}
portal/discovery/relaystate.go
-13
@@ -2,7 +2,6 @@ package discovery
2
3 import (
4 "errors"
5 - "strings"
5 "time"
6
7 "github.com/gosuda/portal-tunnel/v2/types"
@@ -83,15 +82,3 @@ func (state RelayState) discoverable(now time.Time) bool {
82 }
83 return true
84 }
86 -
87 -func (state RelayState) Equal(other RelayState) bool {
88 - stateKey := state.Descriptor.Key()
89 - otherKey := other.Descriptor.Key()
90 - if stateKey != "" && otherKey != "" && stateKey == otherKey {
91 - return true
92 - }
93 -
94 - stateURL := strings.TrimSpace(state.Descriptor.APIHTTPSAddr)
95 - otherURL := strings.TrimSpace(other.Descriptor.APIHTTPSAddr)
96 - return stateURL != "" && otherURL != "" && stateURL == otherURL
97 -}
portal/lease.go
+77
@@ -4,6 +4,8 @@ import (
4 "context"
5 "errors"
6 "fmt"
7 + "net"
8 + "net/url"
9 "strings"
10 "sync"
11 "time"
@@ -292,6 +294,81 @@ func (r *leaseRegistry) countTCPPortLeases() int {
294 return count
295 }
296
297 +func (r *leaseRegistry) snapshot(record *leaseRecord, now time.Time) (types.Lease, bool) {
298 + if record == nil || now.After(record.ExpiresAt) {
299 + return types.Lease{}, false
300 + }
301 +
302 + adminSnapshot := r.AdminSnapshot(record)
303 + since := time.Duration(0)
304 + if !adminSnapshot.LastSeenAt.IsZero() {
305 + since = max(now.Sub(adminSnapshot.LastSeenAt), 0)
306 + }
307 + if adminSnapshot.IsBanned || adminSnapshot.IsDenied || !adminSnapshot.IsApproved || adminSnapshot.Metadata.Hide {
308 + return types.Lease{}, false
309 + }
310 + if adminSnapshot.Ready == 0 && since >= 3*time.Minute {
311 + return types.Lease{}, false
312 + }
313 + return adminSnapshot.Lease, true
314 +}
315 +
316 +func (r *leaseRegistry) LeaseSnapshots(now time.Time) []types.Lease {
317 + r.mu.RLock()
318 + defer r.mu.RUnlock()
319 +
320 + snapshots := make([]types.Lease, 0, len(r.leasesByKey))
321 + for _, record := range r.leasesByKey {
322 + snapshot, ok := r.snapshot(record, now)
323 + if !ok {
324 + continue
325 + }
326 + snapshots = append(snapshots, snapshot)
327 + }
328 + return snapshots
329 +}
330 +
331 +func (r *leaseRegistry) AdminLeaseSnapshots(now time.Time) []types.AdminLease {
332 + r.mu.RLock()
333 + defer r.mu.RUnlock()
334 +
335 + snapshots := make([]types.AdminLease, 0, len(r.leasesByKey))
336 + for _, record := range r.leasesByKey {
337 + if now.After(record.ExpiresAt) {
338 + continue
339 + }
340 + snapshots = append(snapshots, r.AdminSnapshot(record))
341 + }
342 + return snapshots
343 +}
344 +
345 +func (r *leaseRegistry) DiscoverySourceURLs(portalURL string) []string {
346 + baseURL, err := url.Parse(portalURL)
347 + if err != nil {
348 + return nil
349 + }
350 +
351 + r.mu.RLock()
352 + defer r.mu.RUnlock()
353 +
354 + sources := make([]string, 0, len(r.leasesByKey))
355 + for _, record := range r.leasesByKey {
356 + snapshot, ok := r.snapshot(record, time.Now())
357 + if !ok {
358 + continue
359 + }
360 +
361 + sourceURL := *baseURL
362 + sourceURL.Host = snapshot.Hostname
363 + if port := baseURL.Port(); port != "" {
364 + sourceURL.Host = net.JoinHostPort(snapshot.Hostname, port)
365 + }
366 +
367 + sources = append(sources, sourceURL.String())
368 + }
369 + return sources
370 +}
371 +
372 func (r *leaseRegistry) Snapshot(record *leaseRecord) types.Lease {
373 snapshot := types.Lease{
374 Name: record.Name,
portal/server.go
+7 -42
@@ -365,52 +365,17 @@ func (s *Server) PortalURL() string {
365 }
366
367 func (s *Server) LeaseSnapshots() []types.Lease {
368 - s.registry.mu.RLock()
369 - defer s.registry.mu.RUnlock()
370 -
371 - now := time.Now()
372 - records := make([]*leaseRecord, 0, len(s.registry.leasesByKey))
373 - for _, record := range s.registry.leasesByKey {
374 - records = append(records, record)
375 - }
376 - snapshots := make([]types.Lease, 0, len(records))
377 - for _, record := range records {
378 - if now.After(record.ExpiresAt) {
379 - continue
380 - }
381 - adminSnapshot := s.registry.AdminSnapshot(record)
382 - since := time.Duration(0)
383 - if !adminSnapshot.LastSeenAt.IsZero() {
384 - since = max(now.Sub(adminSnapshot.LastSeenAt), 0)
385 - }
386 - if adminSnapshot.IsBanned || adminSnapshot.IsDenied || !adminSnapshot.IsApproved || adminSnapshot.Metadata.Hide {
387 - continue
388 - }
389 - if adminSnapshot.Ready == 0 && since >= 3*time.Minute {
390 - continue
391 - }
392 - snapshots = append(snapshots, adminSnapshot.Lease)
368 + if s == nil || s.registry == nil {
369 + return nil
370 }
394 - return snapshots
371 + return s.registry.LeaseSnapshots(time.Now())
372 }
373
374 func (s *Server) AdminLeaseSnapshots() []types.AdminLease {
398 - s.registry.mu.RLock()
399 - defer s.registry.mu.RUnlock()
400 -
401 - now := time.Now()
402 - records := make([]*leaseRecord, 0, len(s.registry.leasesByKey))
403 - for _, record := range s.registry.leasesByKey {
404 - records = append(records, record)
405 - }
406 - snapshots := make([]types.AdminLease, 0, len(records))
407 - for _, record := range records {
408 - if now.After(record.ExpiresAt) {
409 - continue
410 - }
411 - snapshots = append(snapshots, s.registry.AdminSnapshot(record))
375 + if s == nil || s.registry == nil {
376 + return nil
377 }
413 - return snapshots
378 + return s.registry.AdminLeaseSnapshots(time.Now())
379 }
380
381 func (s *Server) LeaseSnapshotByHostname(hostname string) (types.Lease, bool) {
@@ -639,7 +604,7 @@ func (s *Server) runRelayDiscoveryLoop(ctx context.Context) error {
604 defer ticker.Stop()
605
606 for {
642 - if err := refresher.Refresh(ctx); err != nil {
607 + if err := refresher.Refresh(ctx, s.registry.DiscoverySourceURLs(s.PortalURL())...); err != nil {
608 if ctx.Err() != nil {
609 return nil
610 }
sdk/expose.go
+30 -28
@@ -26,8 +26,8 @@ type Exposure struct {
26 done <-chan struct{}
27
28 identity types.Identity
29 + discovery bool
30 explicitRelays []string
30 - activeRelayURLs []string
31 TargetAddr string
32 UDPAddr string
33 udpEnabled bool
@@ -115,6 +115,7 @@ func Expose(ctx context.Context, cfg ExposeConfig) (*Exposure, error) {
115 cancel: cancel,
116 done: exposureCtx.Done(),
117 identity: identity,
118 + discovery: cfg.Discovery,
119 explicitRelays: append([]string(nil), explicitRelayURLs...),
120 TargetAddr: targetAddr,
121 UDPAddr: udpAddr,
@@ -175,20 +176,12 @@ func (e *Exposure) runDiscoveryLoop(ctx context.Context) {
176 func (e *Exposure) ActiveRelayURLs() []string {
177 e.listenerMu.RLock()
178 defer e.listenerMu.RUnlock()
178 - return append([]string(nil), e.activeRelayURLs...)
179 -}
180 -
181 -func (e *Exposure) clientState() discovery.ClientState {
182 - e.listenerMu.RLock()
183 - defer e.listenerMu.RUnlock()
184 -
185 - return discovery.ClientState{
186 - ActiveRelayURLs: append([]string(nil), e.activeRelayURLs...),
187 - ExplicitRelayURLs: append([]string(nil), e.explicitRelays...),
188 - MaxActiveRelays: e.maxActiveRelays,
189 - RequireUDP: e.udpEnabled,
190 - RequireTCP: e.tcpEnabled,
179 + relayURLs := make([]string, 0, len(e.relayListeners))
180 + for relayURL := range e.relayListeners {
181 + relayURLs = append(relayURLs, relayURL)
182 }
183 + slices.Sort(relayURLs)
184 + return relayURLs
185 }
186
187 func (e *Exposure) Addr() net.Addr {
@@ -277,6 +270,20 @@ func (e *Exposure) WaitDatagramReady(ctx context.Context) ([]string, error) {
270 }
271
272 func (e *Exposure) RunHTTP(ctx context.Context, handler http.Handler, localAddr string) error {
273 + if handler == nil {
274 + handler = http.NotFoundHandler()
275 + }
276 + if e.discovery {
277 + tmp := handler
278 + handler = http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) {
279 + if strings.TrimSpace(r.URL.Path) == types.PathDiscovery {
280 + e.relaySet.ServeDiscovery(w, r)
281 + return
282 + }
283 + tmp.ServeHTTP(w, r)
284 + })
285 + }
286 +
287 e.listenerMu.RLock()
288 hasRelayListeners := len(e.relayListeners) > 0
289 e.listenerMu.RUnlock()
@@ -354,7 +361,6 @@ func (e *Exposure) Close() error {
361 e.listenerMu.Lock()
362 relayListeners := e.relayListeners
363 e.relayListeners = make(map[string]*Listener)
357 - e.activeRelayURLs = nil
364 e.listenerMu.Unlock()
365
366 relayURLs := make([]string, 0, len(relayListeners))
@@ -380,7 +386,15 @@ func (e *Exposure) Close() error {
386 }
387
388 func (e *Exposure) reconcileRelayListeners(failOnError bool) error {
383 - desiredRelayURLs := e.relaySet.PriorityRelays(e.clientState())
389 + clientState := discovery.ClientState{
390 + ActiveRelayURLs: e.ActiveRelayURLs(),
391 + ExplicitRelayURLs: append([]string(nil), e.explicitRelays...),
392 + MaxActiveRelays: e.maxActiveRelays,
393 + RequireUDP: e.udpEnabled,
394 + RequireTCP: e.tcpEnabled,
395 + }
396 +
397 + desiredRelayURLs := e.relaySet.PriorityRelays(clientState)
398
399 e.listenerMu.Lock()
400 staleRelayListeners := make(map[string]*Listener)
@@ -458,15 +472,6 @@ func (e *Exposure) reconcileRelayListeners(failOnError bool) error {
472 go e.runListenerAcceptLoop(listener)
473 }
474
461 - e.listenerMu.Lock()
462 - e.activeRelayURLs = e.activeRelayURLs[:0]
463 - for _, relayURL := range desiredRelayURLs {
464 - if _, ok := e.relayListeners[relayURL]; !ok {
465 - continue
466 - }
467 - e.activeRelayURLs = append(e.activeRelayURLs, relayURL)
468 - }
469 - e.listenerMu.Unlock()
475 if len(removedRelayURLs) > 0 || len(addedRelayURLs) > 0 {
476 log.Info().
477 Strs("added_relays", addedRelayURLs).
@@ -517,9 +522,6 @@ func (e *Exposure) runListenerAcceptLoop(listener *Listener) {
522 if current, ok := e.relayListeners[relayURL]; ok && current == listener {
523 delete(e.relayListeners, relayURL)
524 }
520 - if index := slices.Index(e.activeRelayURLs, relayURL); index >= 0 {
521 - e.activeRelayURLs = slices.Delete(e.activeRelayURLs, index, index+1)
522 - }
525 e.listenerMu.Unlock()
526 }()
527
sdk/expose_test.go
+2 -3
@@ -33,9 +33,8 @@ func TestExposureReconcileRemovesBannedRelayFromActiveSet(t *testing.T) {
33 }
34
35 exposure := &Exposure{
36 - relaySet: mustRelaySet(t, relayA, relayB),
37 - relayListeners: make(map[string]*Listener, 2),
38 - activeRelayURLs: []string{relayA, relayB},
36 + relaySet: mustRelaySet(t, relayA, relayB),
37 + relayListeners: make(map[string]*Listener, 2),
38 }
39 relayAClosed := make(chan struct{})
40 exposure.relayListeners = map[string]*Listener{
types/api.go
+1 -2
@@ -91,8 +91,7 @@ type RegisterResponse struct {
91 type DiscoveryResponse struct {
92 ProtocolVersion string `json:"protocol_version"`
93 GeneratedAt time.Time `json:"generated_at"`
94 - Self RelayDescriptor `json:"self"`
95 - Relays []RelayDescriptor `json:"relays,omitempty"`
94 + Relays []RelayDescriptor `json:"relays"`
95 }
96
97 type QUICControlMessage struct {
types/identity.go
+1
@@ -114,6 +114,7 @@ type RelayDescriptor struct {
114 WireGuardEndpoint string `json:"wireguard_endpoint,omitempty"`
115 OverlayIPv4 string `json:"overlay_ipv4,omitempty"`
116 OverlayCIDRs []string `json:"overlay_cidrs,omitempty"`
117 + Discovery bool `json:"discovery,omitempty"`
118 SupportsUDP bool `json:"supports_udp,omitempty"`
119 SupportsTCP bool `json:"supports_tcp,omitempty"`
120 SupportsOverlayPeer bool `json:"supports_overlay_peer,omitempty"`