separate wireguard and pepper protocol
Lee Yunjin committed
Apr 9, 2026 at 07:17 UTC
e86f0f41c88ac1058f04afcdb63c387869d6d2a1
16 files changed
+1372
-956
cmd/relay-server/main.go
+4
-11
@@ -34,6 +34,7 @@ type relayServerConfig struct {
34
PortalURL string
35
APIPort int
36
SNIPort int
37
+ DiscoveryPort int
38
MinPort int
39
MaxPort int
40
UDPEnabled bool
@@ -42,9 +43,6 @@ type relayServerConfig struct {
43
Bootstraps string
44
DiscoveryEnabled bool
45
MaxRouting int
45
- OverlayEnabled bool
46
- OverlayMaxHops int
47
- OverlayCongestion float64
46
WireGuardPrivateKey string
47
WireGuardEndpoint string
48
OverlayIPv4 string
@@ -76,6 +74,7 @@ func runServeCommand(args []string) error {
74
utils.StringFlagEnv(fs, &cfg.PortalURL, "portal-url", "https://localhost:4017", "portal base URL", "PORTAL_URL")
75
utils.IntFlagEnv(fs, &cfg.APIPort, "api-port", 4017, utils.ParsePortNumber, "Admin/API server port", "API_PORT")
76
utils.IntFlagEnv(fs, &cfg.SNIPort, "sni-port", 443, utils.ParsePortNumber, "TCP SNI router port number", "SNI_PORT")
77
+ utils.IntFlagEnv(fs, &cfg.DiscoveryPort, "discovery-port", 0, utils.ParseOptionalPortNumber, "wireguard overlay listen port (0 uses default)", "DISCOVERY_PORT")
78
utils.IntFlagEnv(fs, &cfg.MinPort, "min-port", 0, utils.ParseOptionalPortNumber, "inclusive minimum lease port shared by UDP and raw TCP transports (0=disabled)", "MIN_PORT")
79
utils.IntFlagEnv(fs, &cfg.MaxPort, "max-port", 0, utils.ParseOptionalPortNumber, "inclusive maximum lease port shared by UDP and raw TCP transports (0=disabled)", "MAX_PORT")
80
utils.BoolFlagEnv(fs, &cfg.UDPEnabled, "udp-enabled", false, "enable UDP relay transport; requires a valid --min-port/--max-port range", "UDP_ENABLED")
@@ -84,9 +83,6 @@ func runServeCommand(args []string) error {
83
utils.StringFlagEnv(fs, &cfg.Bootstraps, "bootstraps", "", "additional bootstrap relay API URLs used for discovery expansion", "BOOTSTRAPS")
84
utils.BoolFlagEnv(fs, &cfg.DiscoveryEnabled, "discovery", false, "serve relay discovery endpoints and poll discovery peers", "DISCOVERY")
85
utils.IntFlagEnv(fs, &cfg.MaxRouting, "max-routing", 1, nil, "maximum number of discovery routing attempts per refresh", "MAX_ROUTING")
87
- utils.BoolFlagEnv(fs, &cfg.OverlayEnabled, "overlay-enabled", false, "enable experimental Pepper overlay route planning", "OVERLAY_ENABLED")
88
- utils.IntFlagEnv(fs, &cfg.OverlayMaxHops, "overlay-max-hops", 0, nil, "Pepper overlay max hops (0 disables overlay route planning)", "OVERLAY_MAX_HOPS")
89
- utils.Float64FlagEnv(fs, &cfg.OverlayCongestion, "overlay-congestion-latency-ms", 120, nil, "latency threshold in ms to trigger reverse-Siamese overlay route selection", "OVERLAY_CONGESTION_LATENCY_MS")
86
utils.StringFlagEnv(fs, &cfg.WireGuardPrivateKey, "wireguard-private-key", "", "wireguard private key for relay overlay", "WIREGUARD_PRIVATE_KEY")
87
utils.StringFlagEnv(fs, &cfg.WireGuardEndpoint, "wireguard-endpoint", "", "wireguard endpoint (host:port) for relay overlay", "WIREGUARD_ENDPOINT")
88
utils.StringFlagEnv(fs, &cfg.OverlayIPv4, "overlay-ipv4", "", "explicit overlay IPv4 override (auto-derived from public key when unset)", "OVERLAY_IPV4")
@@ -131,13 +127,12 @@ func runServeCommand(args []string) error {
127
Str("portal_url", cfg.PortalURL).
128
Str("identity_path", cfg.IdentityPath).
129
Str("admin_settings_path", cfg.AdminSettingsPath).
130
+ Int("discovery_port", cfg.DiscoveryPort).
131
Int("min_port", cfg.MinPort).
132
Int("max_port", cfg.MaxPort).
133
Bool("landing_page_enabled", cfg.LandingPageEnabled).
134
Bool("discovery_enabled", cfg.DiscoveryEnabled).
135
Int("max_routing", cfg.MaxRouting).
139
- Bool("overlay_enabled", cfg.OverlayEnabled).
140
- Int("overlay_max_hops", cfg.OverlayMaxHops).
136
Str("acme_dns_provider", cfg.ACMEDNSProvider).
137
Bool("ens_gasless_enabled", cfg.ENSGaslessEnabled).
138
Bool("udp_enabled", cfg.UDPEnabled).
@@ -162,6 +157,7 @@ func runServer(ctx context.Context, cfg relayServerConfig) error {
157
IdentityPath: cfg.IdentityPath,
158
Bootstraps: bootstraps,
159
WireGuardPrivateKey: cfg.WireGuardPrivateKey,
160
+ DiscoveryPort: cfg.DiscoveryPort,
161
WireGuardEndpoint: cfg.WireGuardEndpoint,
162
OverlayIPv4: cfg.OverlayIPv4,
163
OverlayCIDRs: overlayCIDRs,
@@ -185,9 +181,6 @@ func runServer(ctx context.Context, cfg relayServerConfig) error {
181
TrustProxyHeaders: cfg.TrustProxyHeaders,
182
DiscoveryEnabled: cfg.DiscoveryEnabled,
183
MaxRouting: cfg.MaxRouting,
188
- OverlayEnabled: cfg.OverlayEnabled,
189
- OverlayMaxHops: cfg.OverlayMaxHops,
190
- OverlayCongestion: cfg.OverlayCongestion,
184
MinPort: cfg.MinPort,
185
MaxPort: cfg.MaxPort,
186
UDPEnabled: cfg.UDPEnabled,
portal/api_server.go
+7
-7
@@ -201,13 +201,13 @@ func (s *Server) handleRelayDiscovery(w http.ResponseWriter, r *http.Request) {
201
ExpiresAt: now.Add(2 * types.DiscoveryPollInterval),
202
APIHTTPSAddr: s.cfg.PortalURL,
203
IngressTLSAddr: ingressAddr,
204
- WireGuardPublicKey: wireGuardField(s.wireGuardPeerPlaneEnabled(), s.cfg.WireGuardPublicKey),
205
- WireGuardEndpoint: wireGuardField(s.wireGuardPeerPlaneEnabled(), s.cfg.WireGuardEndpoint),
206
- OverlayIPv4: wireGuardField(s.wireGuardPeerPlaneEnabled(), s.cfg.OverlayIPv4),
207
- OverlayCIDRs: overlayCIDRsField(s.wireGuardPeerPlaneEnabled(), s.cfg.OverlayCIDRs),
204
+ WireGuardPublicKey: wireGuardField(s.wireGuardOverlayEnabled(), s.cfg.WireGuardPublicKey),
205
+ WireGuardEndpoint: wireGuardField(s.wireGuardOverlayEnabled(), s.cfg.WireGuardEndpoint),
206
+ OverlayIPv4: wireGuardField(s.wireGuardOverlayEnabled(), s.cfg.OverlayIPv4),
207
+ OverlayCIDRs: overlayCIDRsField(s.wireGuardOverlayEnabled(), s.cfg.OverlayCIDRs),
208
SupportsUDP: s.cfg.UDPEnabled && s.quicTunnel != nil,
209
SupportsTCP: s.cfg.TCPEnabled,
210
- SupportsOverlayPeer: s.cfg.OverlayEnabled,
210
+ SupportsOverlayPeer: s.wireGuardOverlayEnabled(),
211
Load: float64(s.loadMgr.ActiveConns()),
212
})
213
if err != nil {
@@ -221,8 +221,8 @@ func (s *Server) handleRelayDiscovery(w http.ResponseWriter, r *http.Request) {
221
Self: self,
222
Relays: nil,
223
}
224
- if s.discoveryMgr != nil {
225
- resp.Relays = s.discoveryMgr.ActiveRelayDescriptors()
224
+ if s.relaySet != nil {
225
+ resp.Relays = s.relaySet.AdvertisedDescriptors()
226
}
227
utils.WriteAPIData(w, http.StatusOK, resp)
228
}
portal/discovery/discovery.go
+16
@@ -223,3 +223,19 @@ func DiscoveryUnavailableStatus(err error) (statusCode int, code string, unavail
223
}
224
return 0, "", false
225
}
226
+
227
+func RequireOverlayRelayDescriptor(desc types.RelayDescriptor) error {
228
+ if !desc.SupportsOverlayPeer {
229
+ return errors.New("descriptor does not support overlay peer")
230
+ }
231
+ if desc.WireGuardPublicKey == "" {
232
+ return errors.New("descriptor wireguard public key is required")
233
+ }
234
+ if desc.WireGuardEndpoint == "" {
235
+ return errors.New("descriptor wireguard endpoint is required")
236
+ }
237
+ if desc.OverlayIPv4 == "" {
238
+ return errors.New("descriptor overlay ipv4 is required")
239
+ }
240
+ return nil
241
+}
portal/discovery/relayset.go
+527
-400
@@ -1,9 +1,9 @@
1
package discovery
2
3
import (
4
- "context"
4
"errors"
5
"net/http"
6
+ "reflect"
7
"sort"
8
"strings"
9
"sync"
@@ -15,244 +15,238 @@ import (
15
"github.com/gosuda/portal-tunnel/v2/utils"
16
)
17
18
+type RelayView struct {
19
+ Descriptor types.RelayDescriptor
20
+ FirstSeenAt time.Time
21
+ LastSeenAt time.Time
22
+}
23
+
24
type RelayLocalState struct {
25
Banned bool
26
+ BanReason string
27
+ Bootstrap bool
28
Advertised bool
29
Expired bool
30
+ Reachable bool
31
ConsecutiveFailures int
32
+ LastSuccessAt time.Time
33
+ LastFailureAt time.Time
34
}
35
25
-const discoveryRecoveryFailures = 3
36
+type RelaySummary struct {
37
+ Known int
38
+ Banned int
39
+ Bootstrap int
40
+ Advertised int
41
+ Expired int
42
+ Syncable int
43
+ Reachable int
44
+ Unreachable int
45
+}
46
27
-// RelaySet owns the current discovery state:
28
-// explicit bootstrap relay URLs, the latest descriptor seen for each relay,
29
-// and the minimal local state required to ban or expire a relay.
47
+// RelaySet owns the shared relay discovery view: known relay URLs, stable relay
48
+// id/url mappings, the latest validated descriptor seen for each relay, and common
49
+// process-local relay state such as ban/reachability/failure tracking.
50
+//
51
+// Runtime-specific policy such as bootstrap classification, relay lifecycle, or
52
+// listener ownership belongs in the caller's projection.
53
type RelaySet struct {
31
- mu sync.RWMutex
32
- bootstrapRelayURLs []string
33
- relayKeysByURL map[string]string
34
- relays map[string]types.RelayDescriptor
35
- localByURL map[string]*RelayLocalState
36
- activeRelayURLs []string
37
- activeRelays []types.RelayDescriptor
38
- selfRelayKey string
39
- selfRelayURL string
54
+ mu sync.RWMutex
55
+ knownRelayURLs []string
56
+ relayKeysByURL map[string]string
57
+ relays map[string]RelayView
58
+ localByURL map[string]RelayLocalState
59
+ lastStatusReachable map[string]bool
60
+ lastStatusSummary RelaySummary
61
+ haveLastStatus bool
62
}
63
64
func NewRelaySet() *RelaySet {
65
return &RelaySet{
66
relayKeysByURL: make(map[string]string),
45
- relays: make(map[string]types.RelayDescriptor),
46
- localByURL: make(map[string]*RelayLocalState),
67
+ relays: make(map[string]RelayView),
68
+ localByURL: make(map[string]RelayLocalState),
69
}
70
}
71
50
-func (s *RelaySet) getOrCreateLocalState(relayURL string) *RelayLocalState {
51
- state := s.localByURL[relayURL]
52
- if state == nil {
53
- state = &RelayLocalState{}
54
- s.localByURL[relayURL] = state
55
- }
56
- return state
57
-}
58
-
59
-func relayExpiredAt(desc types.RelayDescriptor, state *RelayLocalState, now time.Time) bool {
60
- if state.Expired {
61
- return true
62
- }
63
- if desc.ExpiresAt.IsZero() {
64
- return false
65
- }
66
- if now.IsZero() {
67
- now = time.Now().UTC()
68
- }
69
- return !desc.ExpiresAt.After(now)
70
-}
71
-
72
-func (s *RelaySet) bootstrapRelayURLSetLocked() map[string]struct{} {
73
- if len(s.bootstrapRelayURLs) == 0 {
72
+func (s *RelaySet) trackedRelayURLs() []string {
73
+ if s == nil {
74
return nil
75
}
76
- out := make(map[string]struct{}, len(s.bootstrapRelayURLs))
77
- for _, relayURL := range s.bootstrapRelayURLs {
78
- out[relayURL] = struct{}{}
79
- }
80
- return out
81
-}
76
83
-func (s *RelaySet) descriptorByURLLocked(relayURL string) (types.RelayDescriptor, bool) {
84
- relayKey, ok := s.relayKeysByURL[relayURL]
85
- if !ok {
86
- return types.RelayDescriptor{}, false
77
+ urls := make([]string, 0, len(s.knownRelayURLs)+len(s.relays))
78
+ seen := make(map[string]struct{}, len(s.knownRelayURLs)+len(s.relays))
79
+ for _, relayURL := range s.knownRelayURLs {
80
+ relayURL = strings.TrimSpace(relayURL)
81
+ if relayURL == "" {
82
+ continue
83
+ }
84
+ if _, ok := seen[relayURL]; ok {
85
+ continue
86
+ }
87
+ seen[relayURL] = struct{}{}
88
+ urls = append(urls, relayURL)
89
}
88
- desc, ok := s.relays[relayKey]
89
- return desc, ok
90
-}
91
-
92
-func (s *RelaySet) isSelfRelayURLLocked(relayURL string) bool {
93
- if s == nil {
94
- return false
90
+ for _, view := range s.relays {
91
+ relayURL := view.Descriptor.APIHTTPSAddr
92
+ if relayURL == "" {
93
+ continue
94
+ }
95
+ if _, ok := seen[relayURL]; ok {
96
+ continue
97
+ }
98
+ seen[relayURL] = struct{}{}
99
+ urls = append(urls, relayURL)
100
}
96
- relayURL = strings.TrimSpace(relayURL)
97
- return relayURL != "" && s.selfRelayURL != "" && relayURL == s.selfRelayURL
101
+ return urls
102
}
103
100
-func (s *RelaySet) isSelfRelayDescriptorLocked(desc types.RelayDescriptor) bool {
104
+func (s *RelaySet) ActiveRelayURLs() []string {
105
if s == nil {
102
- return false
106
+ return nil
107
}
104
- if s.selfRelayKey != "" {
105
- if relayKey := desc.Key(); relayKey != "" && relayKey == s.selfRelayKey {
106
- return true
107
- }
108
+ s.mu.RLock()
109
+ defer s.mu.RUnlock()
110
+ if len(s.knownRelayURLs) == 0 {
111
+ return nil
112
}
109
- return s.isSelfRelayURLLocked(desc.APIHTTPSAddr)
110
-}
113
112
-func (s *RelaySet) removeSelfRelayLocked() {
113
- if s == nil {
114
- return
115
- }
116
- if s.selfRelayURL != "" {
117
- filtered := s.bootstrapRelayURLs[:0]
118
- for _, relayURL := range s.bootstrapRelayURLs {
119
- if s.isSelfRelayURLLocked(relayURL) {
120
- continue
121
- }
122
- filtered = append(filtered, relayURL)
114
+ out := make([]string, 0, len(s.knownRelayURLs))
115
+ for _, relayURL := range s.knownRelayURLs {
116
+ if state, ok := s.localByURL[relayURL]; ok && state.Banned {
117
+ continue
118
}
124
- s.bootstrapRelayURLs = filtered
125
-
126
- delete(s.localByURL, s.selfRelayURL)
127
- delete(s.relayKeysByURL, s.selfRelayURL)
119
+ out = append(out, relayURL)
120
}
129
- if s.selfRelayKey != "" {
130
- if desc, ok := s.relays[s.selfRelayKey]; ok {
131
- delete(s.localByURL, desc.APIHTTPSAddr)
132
- delete(s.relayKeysByURL, desc.APIHTTPSAddr)
133
- }
134
- delete(s.relays, s.selfRelayKey)
121
+ if len(out) == 0 {
122
+ return nil
123
}
124
+ return out
125
}
126
138
-func (s *RelaySet) SetSelfRelay(identity types.Identity, relayURL string) error {
139
- if s == nil {
140
- return nil
127
+func relayExpiredAt(view RelayView, state RelayLocalState, now time.Time) bool {
128
+ if state.Expired {
129
+ return true
130
}
142
- relayURL = strings.TrimSpace(relayURL)
143
- if relayURL != "" {
144
- normalized, err := utils.NormalizeRelayURL(relayURL)
145
- if err != nil {
146
- return err
147
- }
148
- relayURL = normalized
131
+ if view.Descriptor.ExpiresAt.IsZero() {
132
+ return false
133
}
150
-
151
- s.mu.Lock()
152
- defer s.mu.Unlock()
153
- s.selfRelayKey = identity.Key()
154
- s.selfRelayURL = relayURL
155
- s.removeSelfRelayLocked()
156
- s.syncActiveLocked(time.Now().UTC())
157
- return nil
158
-}
159
-
160
-func (s *RelaySet) syncActiveLocked(now time.Time) {
134
if now.IsZero() {
135
now = time.Now().UTC()
136
}
137
+ return !view.Descriptor.ExpiresAt.After(now)
138
+}
139
165
- activeRelayURLs := make([]string, 0, len(s.bootstrapRelayURLs)+len(s.relays))
166
- activeRelays := make([]types.RelayDescriptor, 0, len(s.relays))
167
- seen := make(map[string]struct{}, len(s.bootstrapRelayURLs)+len(s.relays))
168
- for _, relayURL := range s.bootstrapRelayURLs {
169
- if s.isSelfRelayURLLocked(relayURL) {
170
- continue
171
- }
172
- state := s.getOrCreateLocalState(relayURL)
140
+func (s *RelaySet) logStatusChange() {
141
+ now := time.Now().UTC()
142
+ var currentReachable map[string]bool
143
+ trackedRelayURLs := s.trackedRelayURLs()
144
+ if len(trackedRelayURLs) > 0 {
145
+ currentReachable = make(map[string]bool, len(trackedRelayURLs))
146
+ for _, relayURL := range trackedRelayURLs {
147
+ state := s.localByURL[relayURL]
148
+ currentReachable[relayURL] = !state.Banned && state.Reachable
149
+ }
150
+ }
151
+ summary := RelaySummary{}
152
+ for _, relayURL := range trackedRelayURLs {
153
+ summary.Known++
154
+ state := s.localByURL[relayURL]
155
+ relayKey := s.relayKeysByURL[relayURL]
156
+ view, ok := s.relays[relayKey]
157
+ expired := ok && relayExpiredAt(view, state, now) || !ok && state.Expired
158
if state.Banned {
159
+ summary.Banned++
160
continue
161
}
176
- seen[relayURL] = struct{}{}
177
- activeRelayURLs = append(activeRelayURLs, relayURL)
178
- desc, ok := s.descriptorByURLLocked(relayURL)
179
- if !ok || desc.APIHTTPSAddr == "" || !state.Advertised || relayExpiredAt(desc, state, now) {
180
- continue
162
+ if state.Bootstrap {
163
+ summary.Bootstrap++
164
+ }
165
+ if state.Advertised && !expired {
166
+ summary.Advertised++
167
+ }
168
+ if expired {
169
+ summary.Expired++
170
}
182
- activeRelays = append(activeRelays, desc)
171
+ if state.Reachable {
172
+ summary.Reachable++
173
+ } else {
174
+ summary.Unreachable++
175
+ }
176
+ if ok && !state.Bootstrap && !expired && view.Descriptor.SupportsOverlayPeer {
177
+ summary.Syncable++
178
+ }
179
+ }
180
+ if s.haveLastStatus && summary == s.lastStatusSummary && reflect.DeepEqual(currentReachable, s.lastStatusReachable) {
181
+ return
182
}
183
185
- discovered := make([]types.RelayDescriptor, 0, len(s.relays))
186
- for _, desc := range s.relays {
187
- if s.isSelfRelayDescriptorLocked(desc) {
188
- continue
189
- }
190
- relayURL := strings.TrimSpace(desc.APIHTTPSAddr)
191
- if relayURL == "" {
184
+ activated := make([]string, 0)
185
+ deactivated := make([]string, 0)
186
+ for relayURL, reachable := range currentReachable {
187
+ if s.lastStatusReachable == nil || s.lastStatusReachable[relayURL] == reachable {
188
continue
189
}
194
- if _, ok := seen[relayURL]; ok {
195
- continue
190
+ if reachable {
191
+ activated = append(activated, relayURL)
192
+ } else {
193
+ deactivated = append(deactivated, relayURL)
194
}
197
- state := s.getOrCreateLocalState(relayURL)
198
- if state.Banned || !state.Advertised || relayExpiredAt(desc, state, now) {
195
+ }
196
+ for relayURL, reachable := range s.lastStatusReachable {
197
+ if _, ok := currentReachable[relayURL]; ok || !reachable {
198
continue
199
}
201
- seen[relayURL] = struct{}{}
202
- discovered = append(discovered, desc)
200
+ deactivated = append(deactivated, relayURL)
201
}
204
- sort.Slice(discovered, func(i, j int) bool {
205
- return discovered[i].APIHTTPSAddr < discovered[j].APIHTTPSAddr
206
- })
207
- for _, desc := range discovered {
208
- activeRelayURLs = append(activeRelayURLs, desc.APIHTTPSAddr)
209
- activeRelays = append(activeRelays, desc)
210
- }
211
-
212
- s.activeRelayURLs = activeRelayURLs
213
- s.activeRelays = activeRelays
214
-}
202
216
-func (s *RelaySet) ActiveRelayURLs() []string {
217
- if s == nil {
218
- return nil
203
+ event := log.Info().
204
+ Int("banned", summary.Banned).
205
+ Int("bootstrap", summary.Bootstrap).
206
+ Int("advertised", summary.Advertised).
207
+ Int("expired", summary.Expired).
208
+ Int("syncable", summary.Syncable).
209
+ Int("reachable", summary.Reachable).
210
+ Int("unreachable", summary.Unreachable)
211
+ if len(activated) > 0 {
212
+ event = event.Strs("activated", activated)
213
}
220
- s.mu.RLock()
221
- defer s.mu.RUnlock()
222
- if len(s.activeRelayURLs) == 0 {
223
- return nil
214
+ if len(deactivated) > 0 {
215
+ event = event.Strs("deactivated", deactivated)
216
}
225
- return append([]string(nil), s.activeRelayURLs...)
217
+ event.Msg("relay status")
218
+ s.lastStatusReachable = currentReachable
219
+ s.lastStatusSummary = summary
220
+ s.haveLastStatus = true
221
}
222
228
-func (s *RelaySet) bootstrapDescriptors() []types.RelayDescriptor {
223
+func (s *RelaySet) BootstrapDescriptors() []types.RelayDescriptor {
224
if s == nil {
225
return nil
226
}
227
s.mu.RLock()
228
defer s.mu.RUnlock()
234
- if len(s.bootstrapRelayURLs) == 0 {
229
+ if len(s.knownRelayURLs) == 0 {
230
return nil
231
}
232
238
- out := make([]types.RelayDescriptor, 0, len(s.bootstrapRelayURLs))
239
- for _, relayURL := range s.bootstrapRelayURLs {
240
- if s.isSelfRelayURLLocked(relayURL) {
233
+ out := make([]types.RelayDescriptor, 0, len(s.knownRelayURLs))
234
+ for _, relayURL := range s.knownRelayURLs {
235
+ state, ok := s.localByURL[relayURL]
236
+ if !ok || !state.Bootstrap {
237
continue
238
}
243
- if s.getOrCreateLocalState(relayURL).Banned {
244
- continue
245
- }
246
- if desc, ok := s.descriptorByURLLocked(relayURL); ok && desc.APIHTTPSAddr != "" {
247
- out = append(out, desc)
248
- continue
239
+ if relayKey, ok := s.relayKeysByURL[relayURL]; ok {
240
+ if view, ok := s.relays[relayKey]; ok && view.Descriptor.APIHTTPSAddr != "" {
241
+ out = append(out, view.Descriptor)
242
+ continue
243
+ }
244
}
245
out = append(out, types.RelayDescriptor{
246
Identity: types.Identity{
247
Name: utils.PortalRootHost(relayURL),
248
},
249
APIHTTPSAddr: relayURL,
255
- RelayID: relayURL,
250
Version: 1,
251
})
252
}
@@ -262,19 +256,155 @@ func (s *RelaySet) bootstrapDescriptors() []types.RelayDescriptor {
256
return out
257
}
258
265
-func (s *RelaySet) ActiveRelayDescriptors() []types.RelayDescriptor {
259
+func (s *RelaySet) BanRelayURL(relayURL, reason string) bool {
260
+ if s == nil {
261
+ return false
262
+ }
263
+ s.mu.Lock()
264
+ defer s.mu.Unlock()
265
+ relayURL = strings.TrimSpace(relayURL)
266
+ if relayURL == "" {
267
+ return false
268
+ }
269
+
270
+ state := s.localByURL[relayURL]
271
+ reason = strings.TrimSpace(reason)
272
+ changed := !state.Banned || strings.TrimSpace(state.BanReason) != reason
273
+ state.Banned = true
274
+ state.BanReason = reason
275
+ state.Reachable = false
276
+ s.localByURL[relayURL] = state
277
+ if changed {
278
+ s.logStatusChange()
279
+ }
280
+ return changed
281
+}
282
+
283
+func (s *RelaySet) MarkRelayUnreachable(relayURL string) bool {
284
+ if s == nil {
285
+ return false
286
+ }
287
+ s.mu.Lock()
288
+ defer s.mu.Unlock()
289
+ relayURL = strings.TrimSpace(relayURL)
290
+ if relayURL == "" {
291
+ return false
292
+ }
293
+
294
+ state := s.localByURL[relayURL]
295
+ if state.Banned {
296
+ return false
297
+ }
298
+ if !state.Reachable {
299
+ return false
300
+ }
301
+ state.Reachable = false
302
+ s.localByURL[relayURL] = state
303
+ s.logStatusChange()
304
+ return true
305
+}
306
+
307
+func (s *RelaySet) MarkRelayReachable(relayURL string, now time.Time) bool {
308
+ if s == nil {
309
+ return false
310
+ }
311
+ s.mu.Lock()
312
+ defer s.mu.Unlock()
313
+ relayURL = strings.TrimSpace(relayURL)
314
+ if relayURL == "" {
315
+ return false
316
+ }
317
+ if now.IsZero() {
318
+ now = time.Now().UTC()
319
+ }
320
+
321
+ state := s.localByURL[relayURL]
322
+ changed := !state.Reachable || state.ConsecutiveFailures != 0 || state.LastSuccessAt != now
323
+ state.Reachable = true
324
+ state.ConsecutiveFailures = 0
325
+ state.LastSuccessAt = now
326
+ s.localByURL[relayURL] = state
327
+ if changed {
328
+ s.logStatusChange()
329
+ }
330
+ return changed
331
+}
332
+
333
+func (s *RelaySet) MarkRelayFailure(relayURL string, now time.Time) RelayLocalState {
334
+ if s == nil {
335
+ return RelayLocalState{}
336
+ }
337
+ s.mu.Lock()
338
+ defer s.mu.Unlock()
339
+ relayURL = strings.TrimSpace(relayURL)
340
+ if relayURL == "" {
341
+ return RelayLocalState{}
342
+ }
343
+ if now.IsZero() {
344
+ now = time.Now().UTC()
345
+ }
346
+
347
+ state := s.localByURL[relayURL]
348
+ state.Reachable = false
349
+ state.ConsecutiveFailures++
350
+ state.LastFailureAt = now
351
+ s.localByURL[relayURL] = state
352
+ s.logStatusChange()
353
+ return state
354
+}
355
+
356
+func (s *RelaySet) RecordBootstrapDiscoveryFailure(relayURL string, err error, now time.Time) {
357
+ state := s.MarkRelayFailure(relayURL, now)
358
+ if statusCode, code, unavailable := DiscoveryUnavailableStatus(err); unavailable {
359
+ if state.ConsecutiveFailures > 1 {
360
+ return
361
+ }
362
+ event := log.Info().Str("relay", relayURL)
363
+ if statusCode > 0 {
364
+ event = event.Int("status_code", statusCode)
365
+ }
366
+ if code != "" {
367
+ event = event.Str("code", code)
368
+ }
369
+ event.Msg("bootstrap relay discovery unavailable; peer may have discovery disabled")
370
+ return
371
+ }
372
+
373
+ log.Warn().
374
+ Err(err).
375
+ Str("relay", relayURL).
376
+ Msg("bootstrap relay discovery failed")
377
+}
378
+
379
+func (s *RelaySet) AdvertisedDescriptors() []types.RelayDescriptor {
380
if s == nil {
381
return nil
382
}
383
s.mu.RLock()
384
defer s.mu.RUnlock()
271
- if len(s.activeRelays) == 0 {
385
+ if len(s.relays) == 0 {
386
+ return nil
387
+ }
388
+
389
+ now := time.Now().UTC()
390
+ out := make([]types.RelayDescriptor, 0, len(s.relays))
391
+ for _, view := range s.relays {
392
+ state := s.localByURL[view.Descriptor.APIHTTPSAddr]
393
+ if !state.Advertised || relayExpiredAt(view, state, now) || view.Descriptor.APIHTTPSAddr == "" {
394
+ continue
395
+ }
396
+ out = append(out, view.Descriptor)
397
+ }
398
+ if len(out) == 0 {
399
return nil
400
}
274
- return append([]types.RelayDescriptor(nil), s.activeRelays...)
401
+ sort.Slice(out, func(i, j int) bool {
402
+ return out[i].APIHTTPSAddr < out[j].APIHTTPSAddr
403
+ })
404
+ return out
405
}
406
277
-func (s *RelaySet) confirmableDescriptors() []types.RelayDescriptor {
407
+func (s *RelaySet) SyncableDescriptors() []types.RelayDescriptor {
408
if s == nil {
409
return nil
410
}
@@ -285,24 +415,13 @@ func (s *RelaySet) confirmableDescriptors() []types.RelayDescriptor {
415
}
416
417
now := time.Now().UTC()
288
- bootstrapRelayURLs := s.bootstrapRelayURLSetLocked()
418
out := make([]types.RelayDescriptor, 0, len(s.relays))
290
- for _, desc := range s.relays {
291
- if s.isSelfRelayDescriptorLocked(desc) {
292
- continue
293
- }
294
- relayURL := strings.TrimSpace(desc.APIHTTPSAddr)
295
- if relayURL == "" {
419
+ for _, view := range s.relays {
420
+ state := s.localByURL[view.Descriptor.APIHTTPSAddr]
421
+ if state.Bootstrap || relayExpiredAt(view, state, now) || !view.Descriptor.SupportsOverlayPeer {
422
continue
423
}
298
- if _, ok := bootstrapRelayURLs[relayURL]; ok {
299
- continue
300
- }
301
- state := s.getOrCreateLocalState(relayURL)
302
- if state.Banned || relayExpiredAt(desc, state, now) {
303
- continue
304
- }
305
- out = append(out, desc)
424
+ out = append(out, view.Descriptor)
425
}
426
if len(out) == 0 {
427
return nil
@@ -313,300 +432,308 @@ func (s *RelaySet) confirmableDescriptors() []types.RelayDescriptor {
432
return out
433
}
434
316
-func (s *RelaySet) BanRelayURL(relayURL string) {
435
+func (s *RelaySet) Snapshot() map[string]types.RelayState {
436
if s == nil {
318
- return
319
- }
320
- s.mu.Lock()
321
- defer s.mu.Unlock()
322
-
323
- relayURL = strings.TrimSpace(relayURL)
324
- if relayURL == "" {
325
- return
437
+ return nil
438
}
327
- state := s.getOrCreateLocalState(relayURL)
328
- if state.Banned {
329
- return
439
+ s.mu.RLock()
440
+ defer s.mu.RUnlock()
441
+ if len(s.relays) == 0 {
442
+ return nil
443
}
331
- state.Banned = true
332
- s.syncActiveLocked(time.Now().UTC())
444
+
445
+ now := time.Now().UTC()
446
+ snapshot := make(map[string]types.RelayState, len(s.relays))
447
+ for relayKey, view := range s.relays {
448
+ localState := s.localByURL[view.Descriptor.APIHTTPSAddr]
449
+ snapshot[relayKey] = types.RelayState{
450
+ Descriptor: view.Descriptor,
451
+ Bootstrap: localState.Bootstrap,
452
+ Advertised: localState.Advertised,
453
+ Expired: relayExpiredAt(view, localState, now),
454
+ FirstSeenAt: view.FirstSeenAt,
455
+ LastSeenAt: view.LastSeenAt,
456
+ ConsecutiveFailures: localState.ConsecutiveFailures,
457
+ }
458
+ }
459
+ return snapshot
460
}
461
335
-func (s *RelaySet) SetBootstrapRelayURLs(relayURLs []string) {
462
+func (s *RelaySet) ReplaceKnownRelayURLs(relayURLs []string) {
463
if s == nil {
464
return
465
}
466
s.mu.Lock()
467
defer s.mu.Unlock()
341
-
468
filtered := make([]string, 0, len(relayURLs))
343
- seen := make(map[string]struct{}, len(relayURLs))
469
for _, relayURL := range relayURLs {
470
relayURL = strings.TrimSpace(relayURL)
471
if relayURL == "" {
472
continue
473
}
349
- if s.isSelfRelayURLLocked(relayURL) {
350
- continue
474
+ duplicate := false
475
+ for _, existing := range filtered {
476
+ if existing == relayURL {
477
+ duplicate = true
478
+ break
479
+ }
480
}
352
- if _, ok := seen[relayURL]; ok {
481
+ if duplicate {
482
continue
483
}
355
- seen[relayURL] = struct{}{}
484
filtered = append(filtered, relayURL)
357
- s.getOrCreateLocalState(relayURL)
358
- }
359
-
360
- for _, relayURL := range s.bootstrapRelayURLs {
361
- if _, ok := seen[relayURL]; ok {
362
- continue
363
- }
364
- if _, ok := s.relayKeysByURL[relayURL]; ok {
365
- continue
366
- }
367
- delete(s.localByURL, relayURL)
485
}
369
- s.bootstrapRelayURLs = filtered
370
- s.syncActiveLocked(time.Now().UTC())
486
+ s.knownRelayURLs = append([]string(nil), filtered...)
487
}
488
373
-func (s *RelaySet) registerDescriptorLocked(desc types.RelayDescriptor) error {
489
+func (s *RelaySet) registerDescriptor(desc types.RelayDescriptor, now time.Time) (string, bool, bool, error) {
490
+ if s == nil {
491
+ return "", false, false, nil
492
+ }
493
normalized, err := NormalizeDescriptor(desc)
494
if err != nil {
376
- return err
495
+ return "", false, false, err
496
}
497
relayKey := normalized.Key()
498
if relayKey == "" {
380
- return errors.New("descriptor identity is required")
499
+ return "", false, false, errors.New("descriptor identity is required")
500
}
501
if knownRelayKey, ok := s.relayKeysByURL[normalized.APIHTTPSAddr]; ok && knownRelayKey != relayKey {
383
- return errors.New("descriptor identity does not match known relay url")
502
+ return "", false, false, errors.New("descriptor identity does not match known relay url")
503
}
504
386
- previous := s.relays[relayKey]
387
- previousURL := strings.TrimSpace(previous.APIHTTPSAddr)
505
+ if now.IsZero() {
506
+ now = time.Now().UTC()
507
+ }
508
+
509
+ view, ok := s.relays[relayKey]
510
+ added := !ok
511
+ if !ok {
512
+ view.FirstSeenAt = now
513
+ }
514
+ previousURL := view.Descriptor.APIHTTPSAddr
515
+ previousDescriptor := view.Descriptor
516
+ view.Descriptor = normalized
517
+ view.LastSeenAt = now
518
+ s.relays[relayKey] = view
519
+ s.relayKeysByURL[normalized.APIHTTPSAddr] = relayKey
520
if previousURL != "" && previousURL != normalized.APIHTTPSAddr {
521
delete(s.relayKeysByURL, previousURL)
390
- if _, ok := s.relayKeysByURL[previousURL]; !ok {
391
- keepBootstrapState := false
392
- for _, bootstrapURL := range s.bootstrapRelayURLs {
393
- if bootstrapURL == previousURL {
394
- keepBootstrapState = true
395
- break
396
- }
397
- }
398
- if !keepBootstrapState {
399
- delete(s.localByURL, previousURL)
400
- }
401
- }
522
}
523
404
- s.relays[relayKey] = normalized
405
- s.relayKeysByURL[normalized.APIHTTPSAddr] = relayKey
406
- s.getOrCreateLocalState(normalized.APIHTTPSAddr)
407
- return nil
524
+ changed := added || !reflect.DeepEqual(previousDescriptor, normalized)
525
+ return relayKey, added, changed, nil
526
}
527
410
-func (s *RelaySet) ApplyRelayDiscoveryResponse(targetIdentity types.Identity, targetURL string, resp types.DiscoveryResponse, now time.Time) error {
411
- if s == nil {
528
+func relayDiscoveryURLs(selfDescriptor types.RelayDescriptor, relayDescriptors []types.RelayDescriptor) []string {
529
+ relayURLs := make([]string, 0, 1+len(relayDescriptors))
530
+ if apiURL := selfDescriptor.APIHTTPSAddr; apiURL != "" {
531
+ relayURLs = append(relayURLs, apiURL)
532
+ }
533
+ for _, relayDescriptor := range relayDescriptors {
534
+ if apiURL := relayDescriptor.APIHTTPSAddr; apiURL != "" {
535
+ relayURLs = append(relayURLs, apiURL)
536
+ }
537
+ }
538
+ if len(relayURLs) == 0 {
539
return nil
540
}
541
+ return relayURLs
542
+}
543
+
544
+func (s *RelaySet) applyDiscoveryDescriptors(targetIdentity types.Identity, targetURL string, selfDescriptor types.RelayDescriptor, relayDescriptors []types.RelayDescriptor, now time.Time) (relaySetChanged bool, addedRelayCount int, err error) {
545
+ if s == nil {
546
+ return false, 0, nil
547
+ }
548
if strings.TrimSpace(targetIdentity.Name) == "" && strings.TrimSpace(targetIdentity.Address) == "" {
415
- return errors.New("target relay identity is required")
549
+ return false, 0, errors.New("target relay identity is required")
550
}
417
-
418
- selfDescriptor, relayDescriptors, err := ValidateRelayDiscoveryResponse(resp, now)
419
- if err != nil {
420
- return err
551
+ if now.IsZero() {
552
+ now = time.Now().UTC()
553
}
554
if err := ValidateDescriptorTarget(selfDescriptor, targetIdentity, targetURL); err != nil {
423
- return err
555
+ return false, 0, err
556
}
557
426
- s.mu.Lock()
427
- defer s.mu.Unlock()
558
+ apply := func(desc types.RelayDescriptor, advertise, countAdded bool) error {
559
+ _, added, descriptorChanged, err := s.registerDescriptor(desc, now)
560
+ if err != nil {
561
+ return err
562
+ }
563
+ localState := s.localByURL[desc.APIHTTPSAddr]
564
+ wasAdvertised := localState.Advertised
565
+ wasExpired := localState.Expired
566
+ if advertise {
567
+ localState.Advertised = true
568
+ }
569
+ localState.Expired = false
570
+ s.localByURL[desc.APIHTTPSAddr] = localState
571
429
- if err := s.registerDescriptorLocked(selfDescriptor); err != nil {
430
- return err
572
+ changed := added || descriptorChanged || advertise && !wasAdvertised || wasExpired
573
+ if added && countAdded {
574
+ addedRelayCount++
575
+ }
576
+ if changed {
577
+ relaySetChanged = true
578
+ }
579
+ return nil
580
}
432
- selfState := s.getOrCreateLocalState(selfDescriptor.APIHTTPSAddr)
433
- selfState.Advertised = true
434
- selfState.Expired = false
435
- selfState.ConsecutiveFailures = 0
581
582
+ if err := apply(selfDescriptor, true, false); err != nil {
583
+ return false, 0, err
584
+ }
585
for _, relayDescriptor := range relayDescriptors {
438
- if s.isSelfRelayDescriptorLocked(relayDescriptor) {
439
- continue
440
- }
441
- if err := s.registerDescriptorLocked(relayDescriptor); err != nil {
442
- log.Warn().
443
- Err(err).
444
- Str("relay", relayDescriptor.APIHTTPSAddr).
445
- Msg("skipping conflicting discovery relay hint")
446
- continue
586
+ if err := apply(relayDescriptor, false, true); err != nil {
587
+ return false, 0, err
588
}
448
- state := s.getOrCreateLocalState(relayDescriptor.APIHTTPSAddr)
449
- switch {
450
- case state.Expired:
451
- // Fresh hint re-enables direct confirmation but must not restore
452
- // advertisement until the relay confirms itself again.
453
- state.Advertised = false
454
- state.Expired = false
455
- state.ConsecutiveFailures = 0
456
- case !state.Advertised:
457
- state.Expired = false
458
- state.ConsecutiveFailures = 0
459
- }
460
- }
461
- s.syncActiveLocked(now)
462
- return nil
589
+ }
590
+ state := s.localByURL[selfDescriptor.APIHTTPSAddr]
591
+ state.Reachable = true
592
+ state.ConsecutiveFailures = 0
593
+ state.LastSuccessAt = now
594
+ s.localByURL[selfDescriptor.APIHTTPSAddr] = state
595
+ s.logStatusChange()
596
+ return relaySetChanged, addedRelayCount, nil
597
}
598
465
-func (s *RelaySet) RecordDiscoveryFailure(identity types.Identity, relayURL string, err error) (expired bool, expireReason string, consecutiveFailures int) {
466
- if s == nil {
467
- return false, "", 0
468
- }
469
- relayKey := identity.Key()
470
- if relayKey == "" {
471
- return false, "", 0
599
+func (s *RelaySet) ApplyRelayDiscoveryResponse(targetIdentity types.Identity, targetURL string, resp types.DiscoveryResponse, now time.Time) (relayURLs []string, relaySetChanged bool, addedRelayCount int, warnErr error, err error) {
600
+ selfDescriptor, relayDescriptors, validateErr := ValidateRelayDiscoveryResponse(resp, now)
601
+ warnErr = validateErr
602
+ if selfDescriptor.Key() == "" {
603
+ return nil, false, 0, warnErr, validateErr
604
}
473
-
605
s.mu.Lock()
475
- defer s.mu.Unlock()
476
-
477
- desc, ok := s.relays[relayKey]
478
- if !ok {
479
- return false, "", 0
480
- }
481
- relayURL = strings.TrimSpace(relayURL)
482
- if relayURL == "" || s.relayKeysByURL[relayURL] != relayKey {
483
- relayURL = desc.APIHTTPSAddr
484
- }
485
- if relayURL == "" {
486
- return false, "", 0
606
+ relaySetChanged, addedRelayCount, err = s.applyDiscoveryDescriptors(targetIdentity, targetURL, selfDescriptor, relayDescriptors, now)
607
+ s.mu.Unlock()
608
+ if err != nil {
609
+ return nil, false, 0, warnErr, err
610
}
611
+ return relayDiscoveryURLs(selfDescriptor, relayDescriptors), relaySetChanged, addedRelayCount, warnErr, nil
612
+}
613
489
- state := s.getOrCreateLocalState(relayURL)
490
- state.ConsecutiveFailures++
491
- if !state.Expired && state.ConsecutiveFailures >= discoveryRecoveryFailures {
492
- state.Expired = true
493
- state.Advertised = false
494
- s.syncActiveLocked(time.Now().UTC())
495
- return true, "recovery", state.ConsecutiveFailures
614
+func (s *RelaySet) ApplyOverlayRelayDiscoveryResponse(targetIdentity types.Identity, targetURL string, resp types.DiscoveryResponse, now time.Time) (relayURLs []string, relaySetChanged bool, addedRelayCount int, warnErr error, err error) {
615
+ selfDescriptor, relayDescriptors, validateErr := ValidateRelayDiscoveryResponse(resp, now)
616
+ warnErr = validateErr
617
+ if selfDescriptor.Key() == "" {
618
+ return nil, false, 0, warnErr, validateErr
619
}
497
-
498
- var apiErr *types.APIRequestError
499
- if errors.As(err, &apiErr) &&
500
- (apiErr.StatusCode == http.StatusForbidden ||
501
- apiErr.StatusCode == http.StatusNotFound ||
502
- apiErr.StatusCode == http.StatusGone) {
503
- state.Expired = true
504
- state.Advertised = false
505
- s.syncActiveLocked(time.Now().UTC())
506
- return true, "status", state.ConsecutiveFailures
620
+ if err := RequireOverlayRelayDescriptor(selfDescriptor); err != nil {
621
+ return nil, false, 0, warnErr, err
622
}
508
- return false, "", state.ConsecutiveFailures
509
-}
623
511
-func logBootstrapDiscoveryFailure(relayURL string, err error) {
512
- if statusCode, code, unavailable := DiscoveryUnavailableStatus(err); unavailable {
513
- event := log.Info().Str("relay", relayURL)
514
- if statusCode > 0 {
515
- event = event.Int("status_code", statusCode)
516
- }
517
- if code != "" {
518
- event = event.Str("code", code)
624
+ filteredRelayDescriptors := make([]types.RelayDescriptor, 0, len(relayDescriptors))
625
+ for _, relayDescriptor := range relayDescriptors {
626
+ if err := RequireOverlayRelayDescriptor(relayDescriptor); err != nil {
627
+ if warnErr == nil {
628
+ warnErr = err
629
+ }
630
+ continue
631
}
520
- event.Msg("bootstrap relay discovery unavailable")
521
- return
632
+ filteredRelayDescriptors = append(filteredRelayDescriptors, relayDescriptor)
633
}
634
524
- log.Warn().
525
- Err(err).
526
- Str("relay", relayURL).
527
- Msg("bootstrap relay discovery failed")
528
-}
529
-
530
-func logDirectDiscoveryFailure(relayURL string, err error, expired bool, expireReason string, consecutiveFailures int) {
531
- event := log.Warn().
532
- Err(err).
533
- Str("relay", relayURL)
534
- if expired {
535
- event = event.
536
- Bool("expired", true).
537
- Str("reason", expireReason)
538
- if consecutiveFailures > 0 {
539
- event = event.Int("consecutive_failures", consecutiveFailures)
540
- }
635
+ s.mu.Lock()
636
+ relaySetChanged, addedRelayCount, err = s.applyDiscoveryDescriptors(targetIdentity, targetURL, selfDescriptor, filteredRelayDescriptors, now)
637
+ s.mu.Unlock()
638
+ if err != nil {
639
+ return nil, false, 0, warnErr, err
640
}
542
- event.Msg("direct relay discovery failed")
641
+ return relayDiscoveryURLs(selfDescriptor, filteredRelayDescriptors), relaySetChanged, addedRelayCount, warnErr, nil
642
}
643
545
-func (s *RelaySet) refresh(ctx context.Context, rootCAPEM []byte) {
546
- if s == nil {
547
- return
644
+func (s *RelaySet) RegisterBootstrapRelayURLs(inputs []string) ([]string, error) {
645
+ if s == nil || len(inputs) == 0 {
646
+ return nil, nil
647
}
648
550
- for _, bootstrap := range s.bootstrapDescriptors() {
551
- resp, err := DiscoverRelayDiscovery(ctx, bootstrap.APIHTTPSAddr, rootCAPEM, nil)
552
- if err != nil {
553
- if ctx.Err() != nil {
554
- return
555
- }
556
- logBootstrapDiscoveryFailure(bootstrap.APIHTTPSAddr, err)
557
- continue
558
- }
559
- if err := s.ApplyRelayDiscoveryResponse(bootstrap.Identity, bootstrap.APIHTTPSAddr, resp, time.Now().UTC()); err != nil {
560
- log.Warn().
561
- Err(err).
562
- Str("relay", bootstrap.APIHTTPSAddr).
563
- Msg("bootstrap relay discovery failed")
564
- }
649
+ normalized, err := utils.NormalizeRelayURLs(inputs...)
650
+ if err != nil {
651
+ return nil, err
652
}
566
- if ctx.Err() != nil {
567
- return
653
+ normalized, err = utils.ExcludeLocalRelayURLs(normalized...)
654
+ if err != nil {
655
+ return nil, err
656
+ }
657
+ if len(normalized) == 0 {
658
+ return nil, nil
659
}
660
+ s.mu.Lock()
661
+ defer s.mu.Unlock()
662
570
- for _, relay := range s.confirmableDescriptors() {
571
- resp, err := DiscoverRelayDiscovery(ctx, relay.APIHTTPSAddr, rootCAPEM, nil)
572
- if err != nil {
573
- if ctx.Err() != nil {
574
- return
575
- }
576
- expired, expireReason, consecutiveFailures := s.RecordDiscoveryFailure(relay.Identity, relay.APIHTTPSAddr, err)
577
- logDirectDiscoveryFailure(relay.APIHTTPSAddr, err, expired, expireReason, consecutiveFailures)
663
+ existing := make(map[string]struct{}, len(s.knownRelayURLs))
664
+ for _, relayURL := range s.knownRelayURLs {
665
+ existing[relayURL] = struct{}{}
666
+ }
667
+ added := make([]string, 0, len(normalized))
668
+ for _, relayURL := range normalized {
669
+ if _, ok := existing[relayURL]; ok {
670
continue
671
}
580
- if err := s.ApplyRelayDiscoveryResponse(relay.Identity, relay.APIHTTPSAddr, resp, time.Now().UTC()); err != nil {
581
- expired, expireReason, consecutiveFailures := s.RecordDiscoveryFailure(relay.Identity, relay.APIHTTPSAddr, err)
582
- logDirectDiscoveryFailure(relay.APIHTTPSAddr, err, expired, expireReason, consecutiveFailures)
583
- }
672
+ existing[relayURL] = struct{}{}
673
+ s.knownRelayURLs = append(s.knownRelayURLs, relayURL)
674
+ added = append(added, relayURL)
675
+ }
676
+ for _, relayURL := range normalized {
677
+ state := s.localByURL[relayURL]
678
+ state.Bootstrap = true
679
+ state.Reachable = false
680
+ s.localByURL[relayURL] = state
681
+ }
682
+ s.logStatusChange()
683
+ if len(added) == 0 {
684
+ return nil, nil
685
}
686
+ return added, nil
687
}
688
587
-func (s *RelaySet) RunLoop(ctx context.Context, rootCAPEM []byte, syncRuntime func() error) error {
689
+func (s *RelaySet) RecordDiscoveryFailure(identity types.Identity, relayURL string, err error, recoveryFailures int, now time.Time) (expired bool, expireReason string, consecutiveFailures int) {
690
if s == nil {
589
- return nil
691
+ return false, "", 0
692
+ }
693
+ relayKey := identity.Key()
694
+ if relayKey == "" {
695
+ return false, "", 0
696
+ }
697
+ relayURL = strings.TrimSpace(relayURL)
698
+ if relayURL == "" {
699
+ return false, "", 0
700
+ }
701
+ if now.IsZero() {
702
+ now = time.Now().UTC()
703
}
704
592
- ticker := time.NewTicker(types.DiscoveryPollInterval)
593
- defer ticker.Stop()
705
+ s.mu.Lock()
706
+ defer s.mu.Unlock()
707
595
- for {
596
- s.refresh(ctx, rootCAPEM)
597
- if ctx.Err() != nil {
598
- return nil
599
- }
600
- if syncRuntime != nil {
601
- if err := syncRuntime(); err != nil {
602
- return err
603
- }
604
- }
708
+ view, ok := s.relays[relayKey]
709
+ if !ok {
710
+ return false, "", 0
711
+ }
712
606
- select {
607
- case <-ctx.Done():
608
- return nil
609
- case <-ticker.C:
610
- }
713
+ localState := s.localByURL[relayURL]
714
+ localState.Reachable = false
715
+ localState.ConsecutiveFailures++
716
+ localState.LastFailureAt = now
717
+ s.localByURL[relayURL] = localState
718
+ s.logStatusChange()
719
+ if !localState.Expired && localState.ConsecutiveFailures >= recoveryFailures {
720
+ state := s.localByURL[view.Descriptor.APIHTTPSAddr]
721
+ state.Expired = true
722
+ s.localByURL[view.Descriptor.APIHTTPSAddr] = state
723
+ s.logStatusChange()
724
+ return true, "recovery", localState.ConsecutiveFailures
725
+ }
726
+
727
+ var apiErr *types.APIRequestError
728
+ if errors.As(err, &apiErr) &&
729
+ (apiErr.StatusCode == http.StatusForbidden ||
730
+ apiErr.StatusCode == http.StatusNotFound ||
731
+ apiErr.StatusCode == http.StatusGone) {
732
+ state := s.localByURL[view.Descriptor.APIHTTPSAddr]
733
+ state.Expired = true
734
+ s.localByURL[view.Descriptor.APIHTTPSAddr] = state
735
+ s.logStatusChange()
736
+ return true, "status", localState.ConsecutiveFailures
737
}
738
+ return false, "", localState.ConsecutiveFailures
739
}
portal/policy/load_manager.go
new
+49
@@ -0,0 +1,49 @@
1
+package policy
2
+
3
+import "sync/atomic"
4
+
5
+// LoadManager tracks coarse connection counters for relay diagnostics.
6
+type LoadManager struct {
7
+ active int64
8
+ bytesIn int64
9
+ bytesOut int64
10
+}
11
+
12
+func NewLoadManager() *LoadManager {
13
+ return &LoadManager{}
14
+}
15
+
16
+func (m *LoadManager) ActiveConns() int64 {
17
+ if m == nil {
18
+ return 0
19
+ }
20
+ return atomic.LoadInt64(&m.active)
21
+}
22
+
23
+func (m *LoadManager) RecordConnStart() {
24
+ if m == nil {
25
+ return
26
+ }
27
+ atomic.AddInt64(&m.active, 1)
28
+}
29
+
30
+func (m *LoadManager) RecordConnEnd() {
31
+ if m == nil {
32
+ return
33
+ }
34
+ atomic.AddInt64(&m.active, -1)
35
+}
36
+
37
+func (m *LoadManager) RecordBytesIn(n int64) {
38
+ if m == nil || n <= 0 {
39
+ return
40
+ }
41
+ atomic.AddInt64(&m.bytesIn, n)
42
+}
43
+
44
+func (m *LoadManager) RecordBytesOut(n int64) {
45
+ if m == nil || n <= 0 {
46
+ return
47
+ }
48
+ atomic.AddInt64(&m.bytesOut, n)
49
+}
portal/server.go
+189
-277
@@ -5,7 +5,6 @@ import (
5
"crypto/tls"
6
"errors"
7
"fmt"
8
- "hash/crc32"
8
"io"
9
"net"
10
"net/http"
@@ -21,7 +20,6 @@ import (
20
"github.com/gosuda/portal-tunnel/v2/portal/acme"
21
"github.com/gosuda/portal-tunnel/v2/portal/discovery"
22
"github.com/gosuda/portal-tunnel/v2/portal/keyless"
24
- "github.com/gosuda/portal-tunnel/v2/portal/overlay"
23
"github.com/gosuda/portal-tunnel/v2/portal/policy"
24
"github.com/gosuda/portal-tunnel/v2/portal/transport"
25
"github.com/gosuda/portal-tunnel/v2/portal/wireguard"
@@ -31,12 +29,13 @@ import (
29
)
30
31
const (
34
- defaultLeaseTTL = 30 * time.Second
35
- defaultClaimTimeout = 10 * time.Second
36
- defaultIdleKeepalive = 15 * time.Second
37
- defaultReadyQueueLimit = 8
38
- defaultClientHelloWait = 2 * time.Second
39
- defaultControlBodyLimit = 4 << 20
32
+ defaultLeaseTTL = 30 * time.Second
33
+ defaultClaimTimeout = 10 * time.Second
34
+ defaultIdleKeepalive = 15 * time.Second
35
+ defaultReadyQueueLimit = 8
36
+ defaultClientHelloWait = 2 * time.Second
37
+ defaultControlBodyLimit = 4 << 20
38
+ defaultWGRecoveryFailures = 3
39
)
40
41
type ServerConfig struct {
@@ -44,6 +43,7 @@ type ServerConfig struct {
43
IdentityPath string
44
Bootstraps []string
45
WireGuardPrivateKey string
46
+ DiscoveryPort int
47
WireGuardEndpoint string
48
WireGuardPublicKey string
49
OverlayIPv4 string
@@ -57,9 +57,6 @@ type ServerConfig struct {
57
TrustProxyHeaders bool
58
DiscoveryEnabled bool
59
MaxRouting int
60
- OverlayEnabled bool
61
- OverlayMaxHops int
62
- OverlayCongestion float64
60
MinPort int
61
MaxPort int
62
UDPEnabled bool
@@ -71,26 +68,21 @@ type Server struct {
68
sniListener net.Listener
69
apiListener net.Listener
70
apiServer *http.Server
74
- wgPeerListener net.Listener
75
- wgPeerServer *http.Server
71
apiTLSClose io.Closer
72
acmeManager *acme.Manager
73
quicTunnel *quic.Listener
79
- wgRuntime *wireguard.Runtime
74
+ overlay *wireguard.Overlay
75
cancel context.CancelFunc
76
group *errgroup.Group
77
registry *leaseRegistry
78
ports *transport.PortAllocator
79
tcpPorts *transport.PortAllocator
80
loadMgr *policy.LoadManager
86
- weightMgr *policy.WeightManager
81
identity types.Identity
82
cfg ServerConfig
83
trustedProxyCIDRs []*net.IPNet
90
- discoveryMgr *discovery.Manager
91
- overlayPolicy *overlay.RoutePolicy
92
- overlayRoute []uint32
93
- overlayRouteMu sync.RWMutex
84
+ wgConfig wireguard.Config
85
+ relaySet *discovery.RelaySet
86
thumbnails *thumbnail.Service
87
shutdownOnce sync.Once
88
}
@@ -132,52 +124,36 @@ func NewServer(cfg ServerConfig) (*Server, error) {
124
bootstraps = filtered
125
}
126
cfg.Bootstraps = bootstraps
135
- wireGuardConfigured := strings.TrimSpace(cfg.WireGuardPrivateKey) != "" ||
136
- strings.TrimSpace(cfg.WireGuardEndpoint) != "" ||
137
- strings.TrimSpace(cfg.WireGuardPublicKey) != "" ||
138
- strings.TrimSpace(cfg.OverlayIPv4) != "" ||
139
- len(cfg.OverlayCIDRs) > 0
140
- if wireGuardConfigured {
141
- if strings.TrimSpace(cfg.WireGuardPrivateKey) == "" {
142
- return nil, errors.New("wireguard private key is required when overlay is enabled")
143
- }
144
- if strings.TrimSpace(cfg.WireGuardEndpoint) == "" {
145
- return nil, errors.New("wireguard endpoint is required when overlay is enabled")
146
- }
147
- normalizedKey, err := utils.NormalizeWireGuardPrivateKey(cfg.WireGuardPrivateKey)
127
+ generatedWireGuardPrivateKey := ""
128
+ if cfg.DiscoveryEnabled && strings.TrimSpace(cfg.WireGuardPrivateKey) == "" {
129
+ generatedWireGuardPrivateKey, err = utils.GenerateWireGuardPrivateKey()
130
if err != nil {
149
- return nil, fmt.Errorf("normalize wireguard private key: %w", err)
150
- }
151
- cfg.WireGuardPrivateKey = normalizedKey
152
- if strings.TrimSpace(cfg.WireGuardPublicKey) == "" {
153
- publicKey, err := utils.WireGuardPublicKeyFromPrivate(cfg.WireGuardPrivateKey)
154
- if err != nil {
155
- return nil, fmt.Errorf("derive wireguard public key: %w", err)
156
- }
157
- cfg.WireGuardPublicKey = publicKey
158
- }
159
- if strings.TrimSpace(cfg.OverlayIPv4) == "" {
160
- overlayIP, err := utils.DeriveWireGuardOverlayIPv4(cfg.WireGuardPublicKey)
161
- if err != nil {
162
- return nil, fmt.Errorf("derive overlay ipv4: %w", err)
163
- }
164
- cfg.OverlayIPv4 = overlayIP
131
+ return nil, err
132
}
166
- cfg.OverlayCIDRs = utils.NormalizeIPPrefixes(cfg.OverlayCIDRs)
133
+ cfg.WireGuardPrivateKey = generatedWireGuardPrivateKey
134
}
168
- if cfg.OverlayMaxHops < 0 {
169
- return nil, errors.New("overlay max hops must be >= 0")
170
- }
171
- if cfg.OverlayMaxHops > 10 {
172
- return nil, errors.New("overlay max hops must be <= 10")
173
- }
174
- if cfg.OverlayCongestion <= 0 {
175
- cfg.OverlayCongestion = 120
176
- }
177
- cfg.OverlayEnabled = cfg.OverlayEnabled && cfg.OverlayMaxHops > 0
178
- if wireGuardConfigured {
179
- cfg.OverlayEnabled = true
135
+ wgConfig, err := wireguard.NormalizeConfig(rootHost, wireguard.Config{
136
+ PrivateKey: cfg.WireGuardPrivateKey,
137
+ PublicKey: cfg.WireGuardPublicKey,
138
+ Endpoint: cfg.WireGuardEndpoint,
139
+ OverlayIPv4: cfg.OverlayIPv4,
140
+ OverlayCIDRs: cfg.OverlayCIDRs,
141
+ ListenPort: cfg.DiscoveryPort,
142
+ })
143
+ if err != nil {
144
+ return nil, err
145
}
146
+ if generatedWireGuardPrivateKey != "" {
147
+ log.Warn().
148
+ Str("wireguard_public_key", wgConfig.PublicKey).
149
+ Str("wireguard_private_key", generatedWireGuardPrivateKey).
150
+ Msg("generated wireguard private key; set WIREGUARD_PRIVATE_KEY to preserve relay identity")
151
+ }
152
+ cfg.WireGuardPrivateKey = wgConfig.PrivateKey
153
+ cfg.WireGuardPublicKey = wgConfig.PublicKey
154
+ cfg.WireGuardEndpoint = wgConfig.Endpoint
155
+ cfg.OverlayIPv4 = wgConfig.OverlayIPv4
156
+ cfg.OverlayCIDRs = append([]string(nil), wgConfig.OverlayCIDRs...)
157
transportEnabled := cfg.UDPEnabled || cfg.TCPEnabled
158
hasPortRange := cfg.MinPort > 0 && cfg.MaxPort > 0
159
if transportEnabled {
@@ -235,25 +211,17 @@ func NewServer(cfg ServerConfig) (*Server, error) {
211
ports: ports,
212
tcpPorts: tcpPorts,
213
loadMgr: policy.NewLoadManager(),
238
- weightMgr: policy.NewWeightManager(),
214
identity: identity,
215
trustedProxyCIDRs: trustedProxyCIDRs,
216
+ wgConfig: wgConfig,
217
thumbnails: thumbnail.NewService(cfg.HeadlessShellURL),
218
}
243
- if cfg.OverlayEnabled {
244
- s.overlayPolicy = overlay.NewRoutePolicy()
245
- }
219
if cfg.DiscoveryEnabled {
247
- manager, err := discovery.NewManager(discovery.ManagerConfig{
248
- Identity: identity,
249
- PortalURL: cfg.PortalURL,
250
- Bootstraps: cfg.Bootstraps,
251
- MaxRouting: cfg.MaxRouting,
252
- })
253
- if err != nil {
220
+ set := discovery.NewRelaySet()
221
+ if _, err := set.RegisterBootstrapRelayURLs(cfg.Bootstraps); err != nil {
222
return nil, err
223
}
256
- s.discoveryMgr = manager
224
+ s.relaySet = set
225
}
226
227
return s, nil
@@ -302,21 +270,20 @@ func (s *Server) Start(ctx context.Context, apiMux *http.ServeMux) error {
270
s.cancel = cancel
271
s.group = group
272
305
- if s.wireGuardPeerPlaneEnabled() {
306
- if err := s.startWireGuardPeerPlane(); err != nil {
273
+ if s.relaySet != nil && strings.TrimSpace(s.wgConfig.PrivateKey) != "" {
274
+ if err := s.startOverlay(); err != nil {
275
acmeManager.Stop()
276
_ = apiServer.Close()
277
_ = apiCloser.Close()
278
_ = sniListener.Close()
279
cancel()
312
- return fmt.Errorf("start wireguard peer plane: %w", err)
280
+ return err
281
}
282
}
283
284
group.Go(s.runAPIServer)
317
- if s.wgPeerServer != nil && s.wgPeerListener != nil {
318
- group.Go(s.runWireGuardPeerAPIServer)
319
- group.Go(func() error { return s.runWireGuardSyncLoop(groupCtx) })
285
+ if s.overlay != nil {
286
+ group.Go(s.overlay.Serve)
287
}
288
group.Go(func() error { return s.runSNIListener(groupCtx) })
289
group.Go(func() error { return s.runLeaseJanitor(groupCtx, 5*time.Second) })
@@ -345,9 +312,7 @@ func (s *Server) Start(ctx context.Context, apiMux *http.ServeMux) error {
312
Int("min_port", s.cfg.MinPort).
313
Int("max_port", s.cfg.MaxPort).
314
Bool("discovery_enabled", s.cfg.DiscoveryEnabled).
348
- Bool("wireguard_enabled", s.wireGuardPeerPlaneEnabled()).
349
- Bool("overlay_enabled", s.cfg.OverlayEnabled).
350
- Int("overlay_max_hops", s.cfg.OverlayMaxHops).
315
+ Bool("wireguard_enabled", s.wireGuardOverlayEnabled()).
316
Bool("udp_enabled", s.cfg.UDPEnabled).
317
Bool("tcp_enabled", s.cfg.TCPEnabled)
318
if s.quicTunnel != nil {
@@ -413,17 +378,11 @@ func (s *Server) Shutdown(ctx context.Context) error {
378
shutdownErr = err
379
}
380
}
416
- if s.wgPeerServer != nil {
417
- if err := s.wgPeerServer.Shutdown(ctx); err != nil && shutdownErr == nil && !errors.Is(err, http.ErrServerClosed) {
381
+ if s.overlay != nil {
382
+ if err := s.overlay.Shutdown(ctx); err != nil && shutdownErr == nil {
383
shutdownErr = err
384
}
385
}
421
- if s.wgPeerListener != nil {
422
- _ = s.wgPeerListener.Close()
423
- }
424
- if s.wgRuntime != nil {
425
- _ = s.wgRuntime.Close()
426
- }
386
if s.apiTLSClose != nil {
387
_ = s.apiTLSClose.Close()
388
}
@@ -451,14 +410,11 @@ func (s *Server) PortalURL() string {
410
return s.cfg.PortalURL
411
}
412
454
-func (s *Server) wireGuardPeerPlaneEnabled() bool {
413
+func (s *Server) wireGuardOverlayEnabled() bool {
414
if s == nil {
415
return false
416
}
458
- return strings.TrimSpace(s.cfg.WireGuardPrivateKey) != "" &&
459
- strings.TrimSpace(s.cfg.WireGuardEndpoint) != "" &&
460
- strings.TrimSpace(s.cfg.WireGuardPublicKey) != "" &&
461
- strings.TrimSpace(s.cfg.OverlayIPv4) != ""
417
+ return strings.TrimSpace(s.wgConfig.PrivateKey) != ""
418
}
419
420
func (s *Server) LeaseSnapshots() []types.Lease {
@@ -522,129 +478,6 @@ func (s *Server) LeaseSnapshotByHostname(hostname string) (types.Lease, bool) {
478
return s.registry.Snapshot(record), true
479
}
480
525
-func (s *Server) startWireGuardPeerPlane() error {
526
- if s == nil {
527
- return nil
528
- }
529
- runtime, err := wireguard.NewRuntime(wireguard.RuntimeConfig{
530
- PrivateKey: s.cfg.WireGuardPrivateKey,
531
- Endpoint: s.cfg.WireGuardEndpoint,
532
- OverlayIPv4: s.cfg.OverlayIPv4,
533
- })
534
- if err != nil {
535
- return err
536
- }
537
-
538
- listener, err := runtime.ListenTCP(wireguard.DefaultPeerAPIHTTPPort)
539
- if err != nil {
540
- _ = runtime.Close()
541
- return fmt.Errorf("listen wireguard peer api: %w", err)
542
- }
543
-
544
- server := &http.Server{
545
- Handler: s.peerAPIHandler(),
546
- ReadHeaderTimeout: 10 * time.Second,
547
- }
548
-
549
- s.wgRuntime = runtime
550
- s.wgPeerListener = listener
551
- s.wgPeerServer = server
552
- if err := s.syncWireGuardPeers(); err != nil {
553
- _ = server.Close()
554
- _ = runtime.Close()
555
- s.wgRuntime = nil
556
- s.wgPeerListener = nil
557
- s.wgPeerServer = nil
558
- return err
559
- }
560
- return nil
561
-}
562
-
563
-func (s *Server) peerAPIHandler() http.Handler {
564
- mux := http.NewServeMux()
565
- mux.HandleFunc(types.PathRoot, s.handleRoot)
566
- mux.HandleFunc(types.PathHealthz, s.handleHealthz)
567
- if s.cfg.DiscoveryEnabled {
568
- mux.HandleFunc(types.PathDiscovery, s.handleRelayDiscovery)
569
- }
570
- return mux
571
-}
572
-
573
-func (s *Server) runWireGuardPeerAPIServer() error {
574
- if s == nil || s.wgPeerServer == nil || s.wgPeerListener == nil {
575
- return nil
576
- }
577
- err := s.wgPeerServer.Serve(s.wgPeerListener)
578
- if errors.Is(err, http.ErrServerClosed) || errors.Is(err, net.ErrClosed) {
579
- return nil
580
- }
581
- return err
582
-}
583
-
584
-func (s *Server) runWireGuardSyncLoop(ctx context.Context) error {
585
- if s.wgRuntime == nil {
586
- <-ctx.Done()
587
- return nil
588
- }
589
- ticker := time.NewTicker(30 * time.Second)
590
- defer ticker.Stop()
591
- for {
592
- if err := s.syncWireGuardPeers(); err != nil {
593
- log.Warn().Err(err).Msg("sync wireguard peers failed")
594
- }
595
- select {
596
- case <-ctx.Done():
597
- return nil
598
- case <-ticker.C:
599
- }
600
- }
601
-}
602
-
603
-func (s *Server) desiredWireGuardPeers() []types.DesiredPeer {
604
- if s.discoveryMgr == nil {
605
- return nil
606
- }
607
- descs := s.discoveryMgr.ActiveRelayDescriptors()
608
- if len(descs) == 0 {
609
- return nil
610
- }
611
- selfKey := s.identity.Key()
612
- peers := make([]types.DesiredPeer, 0, len(descs))
613
- for _, desc := range descs {
614
- nodeKey := relayNodeKey(desc)
615
- if nodeKey == "" || nodeKey == selfKey {
616
- continue
617
- }
618
- if !desc.SupportsOverlayPeer {
619
- continue
620
- }
621
- if strings.TrimSpace(desc.WireGuardPublicKey) == "" ||
622
- strings.TrimSpace(desc.WireGuardEndpoint) == "" ||
623
- strings.TrimSpace(desc.OverlayIPv4) == "" {
624
- continue
625
- }
626
- allowed := []string{desc.OverlayIPv4 + "/32"}
627
- if len(desc.OverlayCIDRs) > 0 {
628
- allowed = append(allowed, desc.OverlayCIDRs...)
629
- }
630
- peers = append(peers, types.DesiredPeer{
631
- RelayID: nodeKey,
632
- WireGuardPublicKey: desc.WireGuardPublicKey,
633
- WireGuardEndpoint: desc.WireGuardEndpoint,
634
- AllowedIPs: allowed,
635
- })
636
- }
637
- return peers
638
-}
639
-
640
-func (s *Server) syncWireGuardPeers() error {
641
- if s.wgRuntime == nil {
642
- return nil
643
- }
644
- peers := s.desiredWireGuardPeers()
645
- return s.wgRuntime.ApplyPeers(peers)
646
-}
647
-
481
func (s *Server) prepareAPITLS(ctx context.Context) (keyless.TLSMaterialConfig, *acme.Manager, error) {
482
acmeCfg := s.cfg.ACME
483
if baseDomain := utils.NormalizeHostname(acmeCfg.BaseDomain); baseDomain != "" && baseDomain != s.identity.Name {
@@ -828,76 +661,155 @@ func (s *Server) runQUICTunnelListener(listener *quic.Listener) error {
661
}
662
}
663
831
-func (s *Server) runRelayDiscoveryLoop(ctx context.Context) error {
832
- if s.discoveryMgr == nil {
833
- <-ctx.Done()
834
- return nil
664
+func (s *Server) startOverlay() error {
665
+ peerMux := http.NewServeMux()
666
+ peerMux.HandleFunc(types.PathRoot, s.handleRoot)
667
+ peerMux.HandleFunc(types.PathHealthz, s.handleHealthz)
668
+ if s.cfg.DiscoveryEnabled {
669
+ peerMux.HandleFunc(types.PathDiscovery, s.handleRelayDiscovery)
670
}
836
- return s.discoveryMgr.Run(ctx, s.handleDiscoverySnapshot)
837
-}
671
839
-func (s *Server) handleDiscoverySnapshot(_ map[string]types.RelayState) {
840
- if s.discoveryMgr == nil || s.overlayPolicy == nil || !s.cfg.OverlayEnabled {
841
- return
842
- }
843
- descs := s.discoveryMgr.ActiveRelayDescriptors()
844
- if len(descs) == 0 {
845
- s.overlayRouteMu.Lock()
846
- s.overlayRoute = nil
847
- s.overlayRouteMu.Unlock()
848
- return
672
+ overlay, err := wireguard.NewOverlay(s.wgConfig, peerMux)
673
+ if err != nil {
674
+ return fmt.Errorf("start wireguard overlay: %w", err)
675
}
676
851
- candidates := make([]uint32, 0, len(descs))
852
- for _, d := range descs {
853
- nodeKey := relayNodeKey(d)
854
- if nodeKey == "" {
855
- continue
856
- }
857
- candidates = append(candidates, crc32.ChecksumIEEE([]byte(nodeKey)))
858
- }
859
- if len(candidates) == 0 {
860
- return
677
+ if err := overlay.Sync(s.identity.Key(), s.relaySet.Snapshot()); err != nil {
678
+ _ = overlay.Shutdown(context.Background())
679
+ return fmt.Errorf("sync wireguard peers: %w", err)
680
}
862
- selfKey := strings.TrimSpace(s.cfg.PortalURL)
863
- if selfKey == "" {
864
- selfKey = s.identity.Key()
865
- }
866
- if selfKey == "" {
867
- return
868
- }
869
- selfID := crc32.ChecksumIEEE([]byte(selfKey))
870
- route, err := s.overlayPolicy.BuildRouteWithLoad(selfID, candidates, s.cfg.OverlayMaxHops, s.weightMgr.Collect(), s.cfg.OverlayCongestion)
871
- if err != nil {
872
- return
873
- }
874
- s.overlayRouteMu.Lock()
875
- s.overlayRoute = route
876
- s.overlayRouteMu.Unlock()
681
+
682
+ s.overlay = overlay
683
+ return nil
684
}
685
879
-func (s *Server) OverlayRoute() []uint32 {
880
- if s == nil {
881
- return nil
882
- }
883
- s.overlayRouteMu.RLock()
884
- defer s.overlayRouteMu.RUnlock()
885
- if len(s.overlayRoute) == 0 {
686
+func (s *Server) runRelayDiscoveryLoop(ctx context.Context) error {
687
+ if s.relaySet == nil {
688
+ <-ctx.Done()
689
return nil
690
}
888
- out := make([]uint32, len(s.overlayRoute))
889
- copy(out, s.overlayRoute)
890
- return out
891
-}
691
+ ticker := time.NewTicker(types.DiscoveryPollInterval)
692
+ defer ticker.Stop()
693
893
-func relayNodeKey(desc types.RelayDescriptor) string {
894
- if key := strings.TrimSpace(desc.RelayID); key != "" {
895
- return key
896
- }
897
- if key := strings.TrimSpace(desc.APIHTTPSAddr); key != "" {
898
- return key
694
+ for {
695
+ bootstraps := s.relaySet.BootstrapDescriptors()
696
+
697
+ for _, bootstrap := range bootstraps {
698
+ resp, err := discovery.DiscoverRelayDiscovery(ctx, bootstrap.APIHTTPSAddr, nil, nil)
699
+ if err != nil {
700
+ if ctx.Err() != nil {
701
+ return nil
702
+ }
703
+ s.relaySet.RecordBootstrapDiscoveryFailure(bootstrap.APIHTTPSAddr, err, time.Now().UTC())
704
+ continue
705
+ }
706
+
707
+ now := time.Now().UTC()
708
+ var relaySetChanged bool
709
+ var warnErr error
710
+ _, relaySetChanged, _, warnErr, err = s.relaySet.ApplyRelayDiscoveryResponse(bootstrap.Identity, bootstrap.APIHTTPSAddr, resp, now)
711
+ if relaySetChanged && s.overlay != nil {
712
+ if syncErr := s.overlay.Sync(s.identity.Key(), s.relaySet.Snapshot()); syncErr != nil {
713
+ if warnErr == nil {
714
+ warnErr = syncErr
715
+ }
716
+ }
717
+ }
718
+ if err != nil {
719
+ s.relaySet.MarkRelayFailure(bootstrap.APIHTTPSAddr, time.Now().UTC())
720
+ log.Warn().
721
+ Err(err).
722
+ Str("relay", bootstrap.APIHTTPSAddr).
723
+ Msg("bootstrap relay discovery failed")
724
+ continue
725
+ }
726
+ if warnErr != nil {
727
+ log.Warn().
728
+ Err(warnErr).
729
+ Str("relay", bootstrap.APIHTTPSAddr).
730
+ Msg("bootstrap relay discovery completed with warnings")
731
+ }
732
+ }
733
+ if ctx.Err() != nil {
734
+ return nil
735
+ }
736
+
737
+ if s.overlay != nil {
738
+ overlayClient := s.overlay.Client()
739
+ syncableRelays := s.relaySet.SyncableDescriptors()
740
+
741
+ for _, relay := range syncableRelays {
742
+ var failureErr error
743
+
744
+ if err := discovery.RequireOverlayRelayDescriptor(relay); err != nil {
745
+ failureErr = err
746
+ } else {
747
+ discoverURL := "http://" + net.JoinHostPort(relay.OverlayIPv4, fmt.Sprintf("%d", wireguard.DefaultPeerAPIHTTPPort))
748
+ resp, err := discovery.DiscoverRelayDiscovery(ctx, discoverURL, nil, overlayClient)
749
+ if err != nil {
750
+ if ctx.Err() != nil {
751
+ return nil
752
+ }
753
+ failureErr = err
754
+ } else {
755
+ now := time.Now().UTC()
756
+ var relaySetChanged bool
757
+ var warnErr error
758
+ var snapshot map[string]types.RelayState
759
+ _, relaySetChanged, _, warnErr, err = s.relaySet.ApplyOverlayRelayDiscoveryResponse(relay.Identity, relay.APIHTTPSAddr, resp, now)
760
+ if relaySetChanged {
761
+ snapshot = s.relaySet.Snapshot()
762
+ if syncErr := s.overlay.Sync(s.identity.Key(), snapshot); syncErr != nil {
763
+ if warnErr == nil {
764
+ warnErr = syncErr
765
+ }
766
+ }
767
+ }
768
+ if err != nil {
769
+ failureErr = err
770
+ } else {
771
+ if warnErr != nil {
772
+ log.Warn().
773
+ Err(warnErr).
774
+ Str("relay", relay.APIHTTPSAddr).
775
+ Msg("overlay relay discovery completed with warnings")
776
+ }
777
+ continue
778
+ }
779
+ }
780
+ }
781
+
782
+ expired, expireReason, consecutiveFailures := s.relaySet.RecordDiscoveryFailure(relay.Identity, relay.APIHTTPSAddr, failureErr, defaultWGRecoveryFailures, time.Now().UTC())
783
+ if expired {
784
+ if syncErr := s.overlay.Sync(s.identity.Key(), s.relaySet.Snapshot()); syncErr != nil && failureErr == nil {
785
+ failureErr = syncErr
786
+ }
787
+ }
788
+
789
+ event := log.Warn().
790
+ Err(failureErr).
791
+ Str("relay", relay.APIHTTPSAddr)
792
+ if expired {
793
+ event = event.
794
+ Bool("expired", true).
795
+ Str("reason", expireReason)
796
+ if consecutiveFailures > 0 {
797
+ event = event.Int("consecutive_failures", consecutiveFailures)
798
+ }
799
+ }
800
+ event.Msg("overlay relay discovery failed")
801
+ }
802
+ }
803
+ if ctx.Err() != nil {
804
+ return nil
805
+ }
806
+
807
+ select {
808
+ case <-ctx.Done():
809
+ return nil
810
+ case <-ticker.C:
811
+ }
812
}
900
- return desc.Key()
813
}
814
815
func (s *Server) BridgeConns(left, right net.Conn) {
portal/server_test.go
+18
-11
@@ -49,6 +49,11 @@ func mustRelayDescriptor(t *testing.T, relayURL string) types.RelayDescriptor {
49
return desc
50
}
51
52
+func applyRelay(t *testing.T, set *discovery.RelaySet, identity types.Identity, targetURL string, resp types.DiscoveryResponse, now time.Time) error {
53
+ t.Helper()
54
+ return set.ApplyRelayDiscoveryResponseSimple(identity, targetURL, resp, now)
55
+}
56
+
57
func tempIdentityPath(t *testing.T) string {
58
t.Helper()
59
return filepath.Join(t.TempDir(), "relay_identity.json")
@@ -542,7 +547,9 @@ func TestServerDiscoverySkipsSelfRelayHint(t *testing.T) {
547
t.Fatalf("NormalizeDescriptor() self hint error = %v", err)
548
}
549
545
- if err := server.relaySet.ApplyRelayDiscoveryResponse(
550
+ if err := applyRelay(
551
+ t,
552
+ server.relaySet,
553
bootstrapDesc.Identity,
554
bootstrapDesc.APIHTTPSAddr,
555
types.DiscoveryResponse{ProtocolVersion: types.ProtocolVersion, Self: bootstrapDesc, Relays: []types.RelayDescriptor{selfHint}},
@@ -576,7 +583,7 @@ func TestServerRecordVerifiedDiscoveryPeerRequiresDirectConfirmation(t *testing.
583
584
applyDiscovery := func(targetIdentity types.Identity, targetURL string, resp types.DiscoveryResponse) error {
585
now := time.Now().UTC()
579
- return server.relaySet.ApplyRelayDiscoveryResponse(targetIdentity, targetURL, resp, now)
586
+ return applyRelay(t, server.relaySet, targetIdentity, targetURL, resp, now)
587
}
588
589
err = applyDiscovery(
@@ -653,7 +660,7 @@ func TestServerRecordVerifiedDiscoveryPeerExpiresAfterRepeatedDirectFailures(t *
660
relayADesc := mustRelayDescriptor(t, "https://relay-a.example.com")
661
now := time.Now().UTC()
662
656
- if err := server.relaySet.ApplyRelayDiscoveryResponse(
663
+ if err := server.relaySet.ApplyRelayDiscoveryResponseSimple(
664
bootstrapDesc.Identity,
665
bootstrapDesc.APIHTTPSAddr,
666
types.DiscoveryResponse{ProtocolVersion: types.ProtocolVersion, Self: bootstrapDesc, Relays: []types.RelayDescriptor{relayADesc}},
@@ -661,7 +668,7 @@ func TestServerRecordVerifiedDiscoveryPeerExpiresAfterRepeatedDirectFailures(t *
668
); err != nil {
669
t.Fatalf("ApplyRelayDiscoveryResponse() bootstrap error = %v", err)
670
}
664
- if err := server.relaySet.ApplyRelayDiscoveryResponse(
671
+ if err := server.relaySet.ApplyRelayDiscoveryResponseSimple(
672
relayADesc.Identity,
673
relayADesc.APIHTTPSAddr,
674
types.DiscoveryResponse{ProtocolVersion: types.ProtocolVersion, Self: relayADesc},
@@ -718,7 +725,7 @@ func TestServerBootstrapHintDoesNotResetDirectFailureBudget(t *testing.T) {
725
relayADesc := mustRelayDescriptor(t, "https://relay-a.example.com")
726
now := time.Now().UTC()
727
721
- if err := server.relaySet.ApplyRelayDiscoveryResponse(
728
+ if err := server.relaySet.ApplyRelayDiscoveryResponseSimple(
729
bootstrapDesc.Identity,
730
bootstrapDesc.APIHTTPSAddr,
731
types.DiscoveryResponse{ProtocolVersion: types.ProtocolVersion, Self: bootstrapDesc, Relays: []types.RelayDescriptor{relayADesc}},
@@ -726,7 +733,7 @@ func TestServerBootstrapHintDoesNotResetDirectFailureBudget(t *testing.T) {
733
); err != nil {
734
t.Fatalf("ApplyRelayDiscoveryResponse() bootstrap error = %v", err)
735
}
729
- if err := server.relaySet.ApplyRelayDiscoveryResponse(
736
+ if err := server.relaySet.ApplyRelayDiscoveryResponseSimple(
737
relayADesc.Identity,
738
relayADesc.APIHTTPSAddr,
739
types.DiscoveryResponse{ProtocolVersion: types.ProtocolVersion, Self: relayADesc},
@@ -748,7 +755,7 @@ func TestServerBootstrapHintDoesNotResetDirectFailureBudget(t *testing.T) {
755
t.Fatalf("RecordDiscoveryFailure() consecutive = %d, want %d", consecutiveFailures, attempt)
756
}
757
751
- if err := server.relaySet.ApplyRelayDiscoveryResponse(
758
+ if err := server.relaySet.ApplyRelayDiscoveryResponseSimple(
759
bootstrapDesc.Identity,
760
bootstrapDesc.APIHTTPSAddr,
761
types.DiscoveryResponse{ProtocolVersion: types.ProtocolVersion, Self: bootstrapDesc, Relays: []types.RelayDescriptor{relayADesc}},
@@ -801,7 +808,7 @@ func TestServerExpiredDiscoveryPeerNeedsFreshDirectConfirmation(t *testing.T) {
808
relayADesc := mustRelayDescriptor(t, "https://relay-a.example.com")
809
now := time.Now().UTC()
810
804
- if err := server.relaySet.ApplyRelayDiscoveryResponse(
811
+ if err := server.relaySet.ApplyRelayDiscoveryResponseSimple(
812
bootstrapDesc.Identity,
813
bootstrapDesc.APIHTTPSAddr,
814
types.DiscoveryResponse{ProtocolVersion: types.ProtocolVersion, Self: bootstrapDesc, Relays: []types.RelayDescriptor{relayADesc}},
@@ -809,7 +816,7 @@ func TestServerExpiredDiscoveryPeerNeedsFreshDirectConfirmation(t *testing.T) {
816
); err != nil {
817
t.Fatalf("ApplyRelayDiscoveryResponse() bootstrap error = %v", err)
818
}
812
- if err := server.relaySet.ApplyRelayDiscoveryResponse(
819
+ if err := server.relaySet.ApplyRelayDiscoveryResponseSimple(
820
relayADesc.Identity,
821
relayADesc.APIHTTPSAddr,
822
types.DiscoveryResponse{ProtocolVersion: types.ProtocolVersion, Self: relayADesc},
@@ -826,7 +833,7 @@ func TestServerExpiredDiscoveryPeerNeedsFreshDirectConfirmation(t *testing.T) {
833
)
834
}
835
829
- if err := server.relaySet.ApplyRelayDiscoveryResponse(
836
+ if err := server.relaySet.ApplyRelayDiscoveryResponseSimple(
837
bootstrapDesc.Identity,
838
bootstrapDesc.APIHTTPSAddr,
839
types.DiscoveryResponse{ProtocolVersion: types.ProtocolVersion, Self: bootstrapDesc, Relays: []types.RelayDescriptor{relayADesc}},
@@ -848,7 +855,7 @@ func TestServerExpiredDiscoveryPeerNeedsFreshDirectConfirmation(t *testing.T) {
855
t.Fatalf("ActiveRelayDescriptors() = %v, want relay to stay hidden until reconfirmed", advertisedURLs)
856
}
857
851
- if err := server.relaySet.ApplyRelayDiscoveryResponse(
858
+ if err := server.relaySet.ApplyRelayDiscoveryResponseSimple(
859
relayADesc.Identity,
860
relayADesc.APIHTTPSAddr,
861
types.DiscoveryResponse{ProtocolVersion: types.ProtocolVersion, Self: relayADesc},
portal/wireguard/overlay.go
new
+190
@@ -0,0 +1,190 @@
1
+package wireguard
2
+
3
+import (
4
+ "context"
5
+ "errors"
6
+ "fmt"
7
+ "net"
8
+ "net/http"
9
+ "sort"
10
+ "strings"
11
+ "time"
12
+
13
+ "github.com/gosuda/portal-tunnel/v2/types"
14
+ "github.com/gosuda/portal-tunnel/v2/utils"
15
+)
16
+
17
+type Config struct {
18
+ PrivateKey string
19
+ PublicKey string
20
+ Endpoint string
21
+ OverlayIPv4 string
22
+ OverlayCIDRs []string
23
+ ListenPort int
24
+}
25
+
26
+func NormalizeConfig(rootHost string, cfg Config) (Config, error) {
27
+ configured := strings.TrimSpace(cfg.PrivateKey) != "" ||
28
+ strings.TrimSpace(cfg.PublicKey) != "" ||
29
+ strings.TrimSpace(cfg.Endpoint) != "" ||
30
+ strings.TrimSpace(cfg.OverlayIPv4) != "" ||
31
+ len(cfg.OverlayCIDRs) > 0
32
+ if !configured {
33
+ return cfg, nil
34
+ }
35
+
36
+ if strings.TrimSpace(cfg.PrivateKey) == "" {
37
+ return Config{}, errors.New("wireguard private key is required when relay overlay is enabled")
38
+ }
39
+
40
+ privateKey, err := utils.NormalizeWireGuardPrivateKey(cfg.PrivateKey)
41
+ if err != nil {
42
+ return Config{}, fmt.Errorf("normalize wireguard private key: %w", err)
43
+ }
44
+ publicKey, err := utils.WireGuardPublicKeyFromPrivate(privateKey)
45
+ if err != nil {
46
+ return Config{}, fmt.Errorf("derive wireguard public key: %w", err)
47
+ }
48
+ if configuredPublicKey := strings.TrimSpace(cfg.PublicKey); configuredPublicKey != "" && configuredPublicKey != publicKey {
49
+ return Config{}, errors.New("wireguard public key does not match private key")
50
+ }
51
+
52
+ cfg.PrivateKey = privateKey
53
+ cfg.PublicKey = publicKey
54
+ cfg.ListenPort = utils.IntOrDefault(cfg.ListenPort, DefaultListenPort)
55
+ if len(cfg.OverlayCIDRs) > 0 {
56
+ cfg.OverlayCIDRs, err = utils.NormalizeOverlayCIDRs(cfg.OverlayCIDRs)
57
+ if err != nil {
58
+ return Config{}, fmt.Errorf("normalize overlay cidrs: %w", err)
59
+ }
60
+ }
61
+ if strings.TrimSpace(cfg.Endpoint) == "" {
62
+ cfg.Endpoint = net.JoinHostPort(rootHost, fmt.Sprintf("%d", cfg.ListenPort))
63
+ }
64
+ if strings.TrimSpace(cfg.OverlayIPv4) == "" {
65
+ cfg.OverlayIPv4, err = utils.DeriveWireGuardOverlayIPv4(cfg.PublicKey)
66
+ if err != nil {
67
+ return Config{}, fmt.Errorf("derive overlay ipv4: %w", err)
68
+ }
69
+ }
70
+ if err := utils.ValidateWireGuardEndpoint(cfg.Endpoint); err != nil {
71
+ return Config{}, err
72
+ }
73
+ if err := utils.ValidateOverlayIPv4(cfg.OverlayIPv4); err != nil {
74
+ return Config{}, err
75
+ }
76
+ return cfg, nil
77
+}
78
+
79
+type Overlay struct {
80
+ stack *stack
81
+ listener net.Listener
82
+ server *http.Server
83
+}
84
+
85
+func NewOverlay(cfg Config, handler http.Handler) (*Overlay, error) {
86
+ stack, err := newStack(cfg)
87
+ if err != nil {
88
+ return nil, err
89
+ }
90
+
91
+ listener, err := stack.ListenTCP(DefaultPeerAPIHTTPPort)
92
+ if err != nil {
93
+ _ = stack.Close()
94
+ return nil, err
95
+ }
96
+
97
+ server := &http.Server{
98
+ Handler: handler,
99
+ ReadHeaderTimeout: 10 * time.Second,
100
+ }
101
+
102
+ return &Overlay{
103
+ stack: stack,
104
+ listener: listener,
105
+ server: server,
106
+ }, nil
107
+}
108
+
109
+func (o *Overlay) Serve() error {
110
+ if o == nil || o.server == nil || o.listener == nil {
111
+ return nil
112
+ }
113
+
114
+ err := o.server.Serve(o.listener)
115
+ if errors.Is(err, http.ErrServerClosed) || errors.Is(err, net.ErrClosed) {
116
+ return nil
117
+ }
118
+ return err
119
+}
120
+
121
+func (o *Overlay) Shutdown(ctx context.Context) error {
122
+ if o == nil {
123
+ return nil
124
+ }
125
+
126
+ var shutdownErr error
127
+ if o.server != nil {
128
+ err := o.server.Shutdown(ctx)
129
+ if err != nil && !errors.Is(err, http.ErrServerClosed) {
130
+ shutdownErr = errors.Join(shutdownErr, err)
131
+ }
132
+ }
133
+ if o.listener != nil {
134
+ err := o.listener.Close()
135
+ if err != nil && !errors.Is(err, net.ErrClosed) {
136
+ shutdownErr = errors.Join(shutdownErr, err)
137
+ }
138
+ }
139
+ if o.stack != nil {
140
+ shutdownErr = errors.Join(shutdownErr, o.stack.Close())
141
+ }
142
+ return shutdownErr
143
+}
144
+
145
+func (o *Overlay) Client() *http.Client {
146
+ if o == nil || o.stack == nil {
147
+ return nil
148
+ }
149
+ return &http.Client{
150
+ Transport: &http.Transport{
151
+ DialContext: o.stack.DialContext,
152
+ ForceAttemptHTTP2: false,
153
+ },
154
+ }
155
+}
156
+
157
+func (o *Overlay) Sync(selfIdentityKey string, snapshot map[string]types.RelayState) error {
158
+ if o == nil || o.stack == nil {
159
+ return nil
160
+ }
161
+ return o.stack.ApplyPeers(peersForSnapshot(selfIdentityKey, snapshot))
162
+}
163
+
164
+func peersForSnapshot(selfIdentityKey string, snapshot map[string]types.RelayState) []types.DesiredPeer {
165
+ peers := make([]types.DesiredPeer, 0, len(snapshot))
166
+ for _, state := range snapshot {
167
+ if state.Expired {
168
+ continue
169
+ }
170
+ desc := state.Descriptor
171
+ if desc.Key() == selfIdentityKey || !desc.SupportsOverlayPeer {
172
+ continue
173
+ }
174
+ if desc.WireGuardPublicKey == "" || desc.WireGuardEndpoint == "" || desc.OverlayIPv4 == "" {
175
+ continue
176
+ }
177
+
178
+ allowedIPs := []string{desc.OverlayIPv4 + "/32"}
179
+ allowedIPs = append(allowedIPs, desc.OverlayCIDRs...)
180
+ peers = append(peers, types.DesiredPeer{
181
+ WireGuardPublicKey: desc.WireGuardPublicKey,
182
+ WireGuardEndpoint: desc.WireGuardEndpoint,
183
+ AllowedIPs: allowedIPs,
184
+ })
185
+ }
186
+ sort.Slice(peers, func(i, j int) bool {
187
+ return peers[i].WireGuardPublicKey < peers[j].WireGuardPublicKey
188
+ })
189
+ return peers
190
+}
portal/wireguard/runtime.go
deleted
-245
@@ -1,245 +0,0 @@
1
-package wireguard
2
-
3
-import (
4
- "context"
5
- "encoding/json"
6
- "errors"
7
- "fmt"
8
- "net"
9
- "net/http"
10
- "net/netip"
11
- "net/url"
12
- "strconv"
13
- "strings"
14
- "sync"
15
- "time"
16
-
17
- "golang.zx2c4.com/wireguard/conn"
18
- "golang.zx2c4.com/wireguard/device"
19
- "golang.zx2c4.com/wireguard/tun/netstack"
20
-
21
- "github.com/gosuda/portal-tunnel/v2/types"
22
- "github.com/gosuda/portal-tunnel/v2/utils"
23
-)
24
-
25
-const (
26
- DefaultMTU = 1420
27
- DefaultPeerAPIHTTPPort = 7777
28
- DefaultPersistentKeepalive = 25
29
-
30
- defaultPeerRequestTimeout = 15 * time.Second
31
-)
32
-
33
-type RuntimeConfig struct {
34
- PrivateKey string
35
- Endpoint string
36
- OverlayIPv4 string
37
- MTU int
38
-}
39
-
40
-type Runtime struct {
41
- device *device.Device
42
- net *netstack.Net
43
- overlayIP netip.Addr
44
-
45
- mu sync.Mutex
46
- closed bool
47
-}
48
-
49
-func NewRuntime(cfg RuntimeConfig) (*Runtime, error) {
50
- privateKey, err := utils.NormalizeWireGuardPrivateKey(cfg.PrivateKey)
51
- if err != nil {
52
- return nil, fmt.Errorf("normalize wireguard private key: %w", err)
53
- }
54
- listenPort, err := utils.WireGuardListenPort(cfg.Endpoint)
55
- if err != nil {
56
- return nil, err
57
- }
58
-
59
- ip, err := netip.ParseAddr(strings.TrimSpace(cfg.OverlayIPv4))
60
- if err != nil || !ip.Is4() {
61
- return nil, errors.New("overlay ipv4 must be a valid IPv4 address")
62
- }
63
-
64
- mtu := cfg.MTU
65
- if mtu <= 0 {
66
- mtu = DefaultMTU
67
- }
68
-
69
- tunDev, network, err := netstack.CreateNetTUN([]netip.Addr{ip}, nil, mtu)
70
- if err != nil {
71
- return nil, fmt.Errorf("create netstack tun: %w", err)
72
- }
73
-
74
- wgDevice := device.NewDevice(tunDev, conn.NewDefaultBind(), device.NewLogger(device.LogLevelError, "portal-wg"))
75
- privateKeyHex, err := utils.WireGuardKeyHex(privateKey)
76
- if err != nil {
77
- wgDevice.Close()
78
- <-wgDevice.Wait()
79
- return nil, err
80
- }
81
-
82
- config := fmt.Sprintf("private_key=%s\nlisten_port=%d\n", privateKeyHex, listenPort)
83
- if err := wgDevice.IpcSet(config); err != nil {
84
- wgDevice.Close()
85
- <-wgDevice.Wait()
86
- return nil, fmt.Errorf("configure wireguard device: %w", err)
87
- }
88
- if err := wgDevice.Up(); err != nil {
89
- wgDevice.Close()
90
- <-wgDevice.Wait()
91
- return nil, fmt.Errorf("bring wireguard device up: %w", err)
92
- }
93
-
94
- return &Runtime{
95
- device: wgDevice,
96
- net: network,
97
- overlayIP: ip,
98
- }, nil
99
-}
100
-
101
-func (r *Runtime) ListenTCP(port int) (net.Listener, error) {
102
- if r == nil || r.net == nil {
103
- return nil, errors.New("wireguard runtime not initialized")
104
- }
105
- return r.net.ListenTCP(&net.TCPAddr{
106
- IP: net.ParseIP(r.overlayIP.String()),
107
- Port: port,
108
- })
109
-}
110
-
111
-func (r *Runtime) DialContext(ctx context.Context, network, address string) (net.Conn, error) {
112
- if r == nil || r.net == nil {
113
- return nil, errors.New("wireguard runtime not initialized")
114
- }
115
- switch network {
116
- case "tcp", "tcp4", "tcp6":
117
- default:
118
- return nil, fmt.Errorf("unsupported network %q", network)
119
- }
120
-
121
- host, portText, err := net.SplitHostPort(address)
122
- if err != nil {
123
- return nil, err
124
- }
125
- ip, err := netip.ParseAddr(strings.Trim(host, "[]"))
126
- if err != nil {
127
- return nil, err
128
- }
129
- port, err := strconv.Atoi(portText)
130
- if err != nil || port <= 0 || port > 65535 {
131
- return nil, errors.New("invalid tcp port")
132
- }
133
- return r.net.DialContextTCPAddrPort(ctx, netip.AddrPortFrom(ip, uint16(port)))
134
-}
135
-
136
-func (r *Runtime) Discover(ctx context.Context, overlayIPv4 string, port int) (types.DiscoveryResponse, error) {
137
- var resp types.DiscoveryResponse
138
- if err := r.getPeerJSON(ctx, overlayIPv4, port, types.PathDiscovery, &resp); err != nil {
139
- return types.DiscoveryResponse{}, err
140
- }
141
- return resp, nil
142
-}
143
-
144
-func (r *Runtime) getPeerJSON(ctx context.Context, overlayIPv4 string, port int, path string, out any) error {
145
- if r == nil {
146
- return errors.New("wireguard runtime not initialized")
147
- }
148
- if port == 0 {
149
- port = DefaultPeerAPIHTTPPort
150
- }
151
- ip, err := netip.ParseAddr(strings.TrimSpace(overlayIPv4))
152
- if err != nil || !ip.Is4() {
153
- return errors.New("overlay ipv4 must be a valid IPv4 address")
154
- }
155
-
156
- baseURL := &url.URL{
157
- Scheme: "http",
158
- Host: net.JoinHostPort(ip.String(), strconv.Itoa(port)),
159
- Path: path,
160
- }
161
-
162
- httpReq, err := http.NewRequestWithContext(ctx, http.MethodGet, baseURL.String(), nil)
163
- if err != nil {
164
- return err
165
- }
166
-
167
- client := &http.Client{
168
- Transport: &http.Transport{
169
- DialContext: r.DialContext,
170
- ForceAttemptHTTP2: false,
171
- },
172
- Timeout: defaultPeerRequestTimeout,
173
- }
174
-
175
- resp, err := client.Do(httpReq)
176
- if err != nil {
177
- return err
178
- }
179
- defer resp.Body.Close()
180
-
181
- if resp.StatusCode < http.StatusOK || resp.StatusCode >= http.StatusMultipleChoices {
182
- return utils.DecodeAPIRequestError(resp)
183
- }
184
- if out == nil {
185
- return nil
186
- }
187
- return json.NewDecoder(resp.Body).Decode(out)
188
-}
189
-
190
-func (r *Runtime) ApplyPeers(peers []types.DesiredPeer) error {
191
- if r == nil || r.device == nil {
192
- return errors.New("wireguard runtime not initialized")
193
- }
194
-
195
- var builder strings.Builder
196
- builder.WriteString("replace_peers=true\n")
197
-
198
- for _, peer := range peers {
199
- publicKeyHex, err := utils.WireGuardKeyHex(peer.WireGuardPublicKey)
200
- if err != nil {
201
- return fmt.Errorf("normalize peer %q public key: %w", peer.RelayID, err)
202
- }
203
- builder.WriteString("public_key=")
204
- builder.WriteString(publicKeyHex)
205
- builder.WriteByte('\n')
206
- if endpoint := strings.TrimSpace(peer.WireGuardEndpoint); endpoint != "" {
207
- builder.WriteString("endpoint=")
208
- builder.WriteString(endpoint)
209
- builder.WriteByte('\n')
210
- }
211
-
212
- allowedIPs := utils.NormalizeIPPrefixes(peer.AllowedIPs)
213
- for _, allowedIP := range allowedIPs {
214
- builder.WriteString("allowed_ip=")
215
- builder.WriteString(allowedIP)
216
- builder.WriteByte('\n')
217
- }
218
- if DefaultPersistentKeepalive > 0 {
219
- builder.WriteString("persistent_keepalive_interval=")
220
- builder.WriteString(strconv.Itoa(DefaultPersistentKeepalive))
221
- builder.WriteByte('\n')
222
- }
223
- }
224
-
225
- return r.device.IpcSet(builder.String())
226
-}
227
-
228
-func (r *Runtime) Close() error {
229
- if r == nil || r.device == nil {
230
- return nil
231
- }
232
-
233
- r.mu.Lock()
234
- if r.closed {
235
- r.mu.Unlock()
236
- return nil
237
- }
238
- r.closed = true
239
- device := r.device
240
- r.mu.Unlock()
241
-
242
- device.Close()
243
- <-device.Wait()
244
- return nil
245
-}
portal/wireguard/stack.go
new
+248
@@ -0,0 +1,248 @@
1
+package wireguard
2
+
3
+import (
4
+ "context"
5
+ "errors"
6
+ "fmt"
7
+ "net"
8
+ "net/netip"
9
+ "strconv"
10
+ "strings"
11
+ "sync"
12
+ "time"
13
+
14
+ "golang.zx2c4.com/wireguard/conn"
15
+ "golang.zx2c4.com/wireguard/device"
16
+ "golang.zx2c4.com/wireguard/tun/netstack"
17
+
18
+ "github.com/gosuda/portal-tunnel/v2/types"
19
+ "github.com/gosuda/portal-tunnel/v2/utils"
20
+)
21
+
22
+const (
23
+ DefaultMTU = 1420
24
+ DefaultListenPort = 51820
25
+ DefaultPeerAPIHTTPPort = 7777
26
+ DefaultPersistentKeepalive = 25
27
+ defaultEndpointResolveTTL = 3 * time.Second
28
+)
29
+
30
+type stack struct {
31
+ device *device.Device
32
+ net *netstack.Net
33
+ overlayIP netip.Addr
34
+
35
+ mu sync.Mutex
36
+ closed bool
37
+ peerEndpoints map[string]string
38
+}
39
+
40
+func newStack(cfg Config) (*stack, error) {
41
+ canonicalPrivateKey, err := utils.NormalizeWireGuardPrivateKey(cfg.PrivateKey)
42
+ if err != nil {
43
+ return nil, fmt.Errorf("normalize wireguard private key: %w", err)
44
+ }
45
+
46
+ listenPort, err := utils.WireGuardListenPort(cfg.Endpoint)
47
+ if err != nil {
48
+ return nil, err
49
+ }
50
+
51
+ overlayIP, err := netip.ParseAddr(cfg.OverlayIPv4)
52
+ if err != nil || !overlayIP.Is4() {
53
+ return nil, errors.New("overlay ipv4 must be a valid IPv4 address")
54
+ }
55
+
56
+ tunDevice, network, err := netstack.CreateNetTUN([]netip.Addr{overlayIP}, nil, DefaultMTU)
57
+ if err != nil {
58
+ return nil, fmt.Errorf("create netstack tun: %w", err)
59
+ }
60
+
61
+ wgDevice := device.NewDevice(tunDevice, conn.NewDefaultBind(), device.NewLogger(device.LogLevelError, "portal-wg"))
62
+ privateKeyHex, err := utils.WireGuardKeyHex(canonicalPrivateKey)
63
+ if err != nil {
64
+ wgDevice.Close()
65
+ <-wgDevice.Wait()
66
+ return nil, err
67
+ }
68
+
69
+ config := fmt.Sprintf("private_key=%s\nlisten_port=%d\n", privateKeyHex, listenPort)
70
+ if err := wgDevice.IpcSet(config); err != nil {
71
+ wgDevice.Close()
72
+ <-wgDevice.Wait()
73
+ return nil, fmt.Errorf("configure wireguard device: %w", err)
74
+ }
75
+ if err := wgDevice.Up(); err != nil {
76
+ wgDevice.Close()
77
+ <-wgDevice.Wait()
78
+ return nil, fmt.Errorf("bring wireguard device up: %w", err)
79
+ }
80
+
81
+ return &stack{
82
+ device: wgDevice,
83
+ net: network,
84
+ overlayIP: overlayIP,
85
+ peerEndpoints: map[string]string{},
86
+ }, nil
87
+}
88
+
89
+func (s *stack) ListenTCP(port int) (net.Listener, error) {
90
+ if s == nil || s.net == nil {
91
+ return nil, errors.New("wireguard is not initialized")
92
+ }
93
+ return s.net.ListenTCP(&net.TCPAddr{
94
+ IP: net.ParseIP(s.overlayIP.String()),
95
+ Port: port,
96
+ })
97
+}
98
+
99
+func (s *stack) DialContext(ctx context.Context, network, address string) (net.Conn, error) {
100
+ if s == nil || s.net == nil {
101
+ return nil, errors.New("wireguard is not initialized")
102
+ }
103
+ switch network {
104
+ case "tcp", "tcp4", "tcp6":
105
+ default:
106
+ return nil, fmt.Errorf("unsupported network %q", network)
107
+ }
108
+
109
+ host, portText, err := net.SplitHostPort(address)
110
+ if err != nil {
111
+ return nil, err
112
+ }
113
+ ip, err := netip.ParseAddr(strings.Trim(host, "[]"))
114
+ if err != nil {
115
+ return nil, err
116
+ }
117
+ port, err := strconv.Atoi(portText)
118
+ if err != nil || port <= 0 || port > 65535 {
119
+ return nil, errors.New("invalid tcp port")
120
+ }
121
+ return s.net.DialContextTCPAddrPort(ctx, netip.AddrPortFrom(ip, uint16(port)))
122
+}
123
+
124
+func (s *stack) ApplyPeers(peers []types.DesiredPeer) error {
125
+ if s == nil || s.device == nil {
126
+ return errors.New("wireguard is not initialized")
127
+ }
128
+
129
+ var builder strings.Builder
130
+ builder.WriteString("replace_peers=true\n")
131
+ var warnErr error
132
+ nextPeerEndpoints := map[string]string{}
133
+
134
+ for _, peer := range peers {
135
+ peerKey := strings.TrimSpace(peer.WireGuardPublicKey)
136
+ publicKeyHex, err := utils.WireGuardKeyHex(peer.WireGuardPublicKey)
137
+ if err != nil {
138
+ return fmt.Errorf("normalize peer %q public key: %w", peerKey, err)
139
+ }
140
+
141
+ resolvedEndpoint := ""
142
+ if endpoint := peer.WireGuardEndpoint; endpoint != "" {
143
+ resolvedEndpoint, err = resolvePeerEndpoint(endpoint)
144
+ if err != nil {
145
+ s.mu.Lock()
146
+ currentEndpoint := s.peerEndpoints[publicKeyHex]
147
+ s.mu.Unlock()
148
+ if currentEndpoint != "" {
149
+ warnErr = errors.Join(warnErr, fmt.Errorf("resolve peer %q endpoint: %w; using current endpoint %q", peerKey, err, currentEndpoint))
150
+ resolvedEndpoint = currentEndpoint
151
+ } else {
152
+ warnErr = errors.Join(warnErr, fmt.Errorf("resolve peer %q endpoint: %w", peerKey, err))
153
+ continue
154
+ }
155
+ }
156
+ }
157
+
158
+ builder.WriteString("public_key=")
159
+ builder.WriteString(publicKeyHex)
160
+ builder.WriteByte('\n')
161
+ if resolvedEndpoint != "" {
162
+ builder.WriteString("endpoint=")
163
+ builder.WriteString(resolvedEndpoint)
164
+ builder.WriteByte('\n')
165
+ nextPeerEndpoints[publicKeyHex] = resolvedEndpoint
166
+ }
167
+
168
+ allowedIPs := utils.NormalizeIPPrefixes(peer.AllowedIPs)
169
+ for _, allowedIP := range allowedIPs {
170
+ builder.WriteString("allowed_ip=")
171
+ builder.WriteString(allowedIP)
172
+ builder.WriteByte('\n')
173
+ }
174
+ if DefaultPersistentKeepalive > 0 {
175
+ builder.WriteString("persistent_keepalive_interval=")
176
+ builder.WriteString(strconv.Itoa(DefaultPersistentKeepalive))
177
+ builder.WriteByte('\n')
178
+ }
179
+ }
180
+
181
+ if err := s.device.IpcSet(builder.String()); err != nil {
182
+ return err
183
+ }
184
+ s.mu.Lock()
185
+ s.peerEndpoints = nextPeerEndpoints
186
+ s.mu.Unlock()
187
+ return warnErr
188
+}
189
+
190
+func resolvePeerEndpoint(raw string) (string, error) {
191
+ endpoint := strings.TrimSpace(raw)
192
+ if endpoint == "" {
193
+ return "", errors.New("wireguard endpoint is required")
194
+ }
195
+
196
+ host, port, err := net.SplitHostPort(endpoint)
197
+ if err != nil {
198
+ return "", err
199
+ }
200
+
201
+ host = strings.Trim(host, "[]")
202
+ if host == "" {
203
+ return "", errors.New("wireguard endpoint host is required")
204
+ }
205
+
206
+ if ip, err := netip.ParseAddr(host); err == nil {
207
+ return net.JoinHostPort(ip.String(), port), nil
208
+ }
209
+
210
+ ctx, cancel := context.WithTimeout(context.Background(), defaultEndpointResolveTTL)
211
+ defer cancel()
212
+
213
+ addrs, err := net.DefaultResolver.LookupNetIP(ctx, "ip", host)
214
+ if err != nil {
215
+ return "", fmt.Errorf("lookup %q: %w", host, err)
216
+ }
217
+ if len(addrs) == 0 {
218
+ return "", fmt.Errorf("lookup %q: no IP addresses found", host)
219
+ }
220
+
221
+ selected := addrs[0]
222
+ for _, addr := range addrs {
223
+ if addr.Is4() {
224
+ selected = addr
225
+ break
226
+ }
227
+ }
228
+ return net.JoinHostPort(selected.String(), port), nil
229
+}
230
+
231
+func (s *stack) Close() error {
232
+ if s == nil || s.device == nil {
233
+ return nil
234
+ }
235
+
236
+ s.mu.Lock()
237
+ if s.closed {
238
+ s.mu.Unlock()
239
+ return nil
240
+ }
241
+ s.closed = true
242
+ device := s.device
243
+ s.mu.Unlock()
244
+
245
+ device.Close()
246
+ <-device.Wait()
247
+ return nil
248
+}
sdk/expose_test.go
+4
-4
@@ -55,7 +55,7 @@ func TestExposureBanRelayURLMovesRelay(t *testing.T) {
55
relayB: {},
56
}
57
58
- exposure.relaySet.BanRelayURL(relayA)
58
+ exposure.relaySet.BanRelayURL(relayA, "test")
59
exposure.listenerMu.Lock()
60
delete(exposure.relayListeners, relayA)
61
exposure.listenerMu.Unlock()
@@ -86,7 +86,7 @@ func TestExposureSetRelayURLsSkipsBannedRelay(t *testing.T) {
86
relaySet: discovery.NewRelaySet(),
87
relayListeners: make(map[string]*Listener, 1),
88
}
89
- exposure.relaySet.BanRelayURL(relayB)
89
+ exposure.relaySet.BanRelayURL(relayB, "test")
90
exposure.relayListeners = map[string]*Listener{
91
relayA: {},
92
}
@@ -167,12 +167,12 @@ func TestExposurePinDiscoveredDescriptorAllowsURLChangeForSameIdentity(t *testin
167
exposure := &Exposure{relaySet: discovery.NewRelaySet()}
168
desc := mustRelayDescriptor(t, "relay-a", "https://relay-a.example")
169
170
- if err := exposure.relaySet.ApplyRelayDiscoveryResponse(desc.Identity, desc.APIHTTPSAddr, types.DiscoveryResponse{ProtocolVersion: types.ProtocolVersion, Self: desc}, time.Now().UTC()); err != nil {
170
+ if err := exposure.relaySet.ApplyRelayDiscoveryResponseSimple(desc.Identity, desc.APIHTTPSAddr, types.DiscoveryResponse{ProtocolVersion: types.ProtocolVersion, Self: desc}, time.Now().UTC()); err != nil {
171
t.Fatalf("ApplyRelayDiscoveryResponse() error = %v", err)
172
}
173
174
changedURL := mustRelayDescriptor(t, desc.Name, "https://relay-b.example")
175
- err := exposure.relaySet.ApplyRelayDiscoveryResponse(desc.Identity, "", types.DiscoveryResponse{ProtocolVersion: types.ProtocolVersion, Self: changedURL}, time.Now().UTC())
175
+ err := exposure.relaySet.ApplyRelayDiscoveryResponseSimple(desc.Identity, "", types.DiscoveryResponse{ProtocolVersion: types.ProtocolVersion, Self: changedURL}, time.Now().UTC())
176
if err != nil {
177
t.Fatalf("ApplyRelayDiscoveryResponse() error = %v, want nil for same relay identity", err)
178
}
sdk/listener.go
+1
-1
@@ -546,7 +546,7 @@ func (l *Listener) closed() bool {
546
547
func (l *Listener) ban() {
548
if l.relaySet != nil && l.api != nil && l.api.baseURL != nil {
549
- l.relaySet.BanRelayURL(l.api.baseURL.String())
549
+ l.relaySet.BanRelayURL(l.api.baseURL.String(), "manual")
550
}
551
_ = l.Close()
552
}
types/relay_state.go
new
+16
@@ -0,0 +1,16 @@
1
+package types
2
+
3
+import "time"
4
+
5
+// RelayState captures the last-known descriptor and local state for a relay
6
+// observed through discovery.
7
+type RelayState struct {
8
+ Descriptor RelayDescriptor `json:"descriptor"`
9
+ Bootstrap bool `json:"bootstrap,omitempty"`
10
+ Advertised bool `json:"advertised,omitempty"`
11
+ Expired bool `json:"expired,omitempty"`
12
+ Banned bool `json:"banned,omitempty"`
13
+ FirstSeenAt time.Time `json:"first_seen_at"`
14
+ LastSeenAt time.Time `json:"last_seen_at"`
15
+ ConsecutiveFailures int `json:"consecutive_failures,omitempty"`
16
+}
types/types.go
+3
@@ -5,6 +5,9 @@ const (
5
ProtocolVersion = "5"
6
PortalRelayRegistryURL = "https://raw.githubusercontent.com/gosuda/portal-tunnel/main/registry.json"
7
8
+ MinDiscoveryRoutingAttempts = 1
9
+ MaxDiscoveryRoutingAttempts = 32
10
+
11
HeaderAccessToken = "X-Portal-Access-Token"
12
MarkerKeepalive = byte(0x00)
13
MarkerRawStart = byte(0x01)
utils/discovery.go
new
+16
@@ -0,0 +1,16 @@
1
+package utils
2
+
3
+import (
4
+ "fmt"
5
+
6
+ "github.com/gosuda/portal-tunnel/v2/types"
7
+)
8
+
9
+// ValidateMaxRouting ensures the configured MaxRouting is within the supported range.
10
+func ValidateMaxRouting(attempts int) error {
11
+ if attempts < types.MinDiscoveryRoutingAttempts || attempts > types.MaxDiscoveryRoutingAttempts {
12
+ return fmt.Errorf("max routing attempts must be between %d and %d (got %d)",
13
+ types.MinDiscoveryRoutingAttempts, types.MaxDiscoveryRoutingAttempts, attempts)
14
+ }
15
+ return nil
16
+}
utils/wireguard.go
+84
@@ -1,12 +1,15 @@
1
package utils
2
3
import (
4
+ "crypto/rand"
5
"crypto/sha256"
6
"encoding/base64"
7
"encoding/hex"
8
"errors"
9
+ "fmt"
10
"net"
11
"net/netip"
12
+ "sort"
13
"strconv"
14
"strings"
15
@@ -22,6 +25,15 @@ func NormalizeWireGuardPrivateKey(raw string) (string, error) {
25
return base64.StdEncoding.EncodeToString(key[:]), nil
26
}
27
28
+func GenerateWireGuardPrivateKey() (string, error) {
29
+ var key [32]byte
30
+ if _, err := rand.Read(key[:]); err != nil {
31
+ return "", err
32
+ }
33
+ clampWireGuardPrivateKey(&key)
34
+ return base64.StdEncoding.EncodeToString(key[:]), nil
35
+}
36
+
37
func WireGuardPublicKeyFromPrivate(raw string) (string, error) {
38
privateKey, err := decodeWireGuardKey(raw)
39
if err != nil {
@@ -103,3 +115,75 @@ func clampWireGuardPrivateKey(key *[32]byte) {
115
key[0] &= 248
116
key[31] = (key[31] & 127) | 64
117
}
118
+
119
+func ValidateWireGuardPublicKey(raw string) error {
120
+ key := strings.TrimSpace(raw)
121
+ if key == "" {
122
+ return errors.New("wireguard_public_key is required")
123
+ }
124
+ decoded, err := base64.StdEncoding.DecodeString(key)
125
+ if err != nil {
126
+ return errors.New("wireguard_public_key must be base64 encoded")
127
+ }
128
+ if len(decoded) != 32 {
129
+ return errors.New("wireguard_public_key must be 32 bytes")
130
+ }
131
+ return nil
132
+}
133
+
134
+func ValidateWireGuardEndpoint(raw string) error {
135
+ endpoint := strings.TrimSpace(raw)
136
+ if endpoint == "" {
137
+ return errors.New("wireguard_endpoint is required")
138
+ }
139
+ host, port, err := net.SplitHostPort(endpoint)
140
+ if err != nil {
141
+ return errors.New("wireguard_endpoint must be host:port")
142
+ }
143
+ if strings.TrimSpace(host) == "" {
144
+ return errors.New("wireguard_endpoint host is required")
145
+ }
146
+ portNum, err := strconv.Atoi(port)
147
+ if err != nil || portNum <= 0 || portNum > 65535 {
148
+ return errors.New("wireguard_endpoint port is invalid")
149
+ }
150
+ return nil
151
+}
152
+
153
+func ValidateOverlayIPv4(raw string) error {
154
+ ipText := strings.TrimSpace(raw)
155
+ if ipText == "" {
156
+ return errors.New("overlay_ipv4 is required")
157
+ }
158
+ ip := net.ParseIP(ipText)
159
+ if ip == nil || ip.To4() == nil {
160
+ return errors.New("overlay_ipv4 must be a valid IPv4 address")
161
+ }
162
+ return nil
163
+}
164
+
165
+func NormalizeOverlayCIDRs(inputs []string) ([]string, error) {
166
+ if len(inputs) == 0 {
167
+ return nil, nil
168
+ }
169
+ seen := make(map[string]struct{}, len(inputs))
170
+ out := make([]string, 0, len(inputs))
171
+ for _, input := range inputs {
172
+ input = strings.TrimSpace(input)
173
+ if input == "" {
174
+ continue
175
+ }
176
+ _, network, err := net.ParseCIDR(input)
177
+ if err != nil {
178
+ return nil, fmt.Errorf("invalid overlay cidr %q", input)
179
+ }
180
+ normalized := network.String()
181
+ if _, ok := seen[normalized]; ok {
182
+ continue
183
+ }
184
+ seen[normalized] = struct{}{}
185
+ out = append(out, normalized)
186
+ }
187
+ sort.Strings(out)
188
+ return out, nil
189
+}