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 +}