refact: add relayset

rabbitprincess committed Mar 30, 2026 at 00:01 UTC 1c692c2c53d7e7ab374e33eb85d2ae639b355df5
17 files changed +1437 -1403
frontend/src/components/ServerListView.tsx
+3 -2
@@ -364,10 +364,11 @@ export function ServerListView({
364 };
365 }, [isAdmin, officialRegistryRelays]);
366
367 + const officialRegistryList = officialRegistryRelays ?? [];
368 const isAllSelected =
369 allLeaseIds.length > 0 &&
370 allLeaseIds.every((id) => selectedLeaseIds.has(id));
370 - const officialRegistryAvailable = (officialRegistryRelays?.length ?? 0) > 0;
371 + const officialRegistryAvailable = officialRegistryList.length > 0;
372
373 const handleSelectAll = () => {
374 if (isAllSelected) {
@@ -809,7 +810,7 @@ export function ServerListView({
810 <div className="mt-6 rounded-xl border border-border/80 bg-secondary/35 p-5 sm:p-6">
811 {officialRegistryAvailable ? (
812 <div className="grid gap-3 md:grid-cols-2 xl:grid-cols-3">
812 - {officialRegistryRelays.map((relay) => {
813 + {officialRegistryList.map((relay) => {
814 return (
815 <div
816 key={relay.url}
portal/api_server.go
+24 -68
@@ -84,7 +84,7 @@ func (s *Server) apiHandler(base *http.ServeMux, keylessSignerHandler http.Handl
84 base.ServeHTTP(w, r)
85 return
86 }
87 - discovery.ServeHTTP(w, r, s.discover)
87 + s.handleRelayDiscovery(w, r)
88 case types.PathV1Sign:
89 if keylessSignerHandler == nil {
90 http.NotFound(w, r)
@@ -108,68 +108,12 @@ func (s *Server) handleHealthz(w http.ResponseWriter, _ *http.Request) {
108 utils.WriteAPIData(w, http.StatusOK, map[string]any{"status": "ok"})
109 }
110
111 -func (s *Server) discover(_ context.Context, req types.DiscoverRequest) (types.DiscoverResponse, error) {
112 - self, err := s.discoverySelfDescriptor()
113 - if err != nil {
114 - return types.DiscoverResponse{}, err
115 - }
116 - resp := types.DiscoverResponse{
117 - ProtocolVersion: 1,
118 - GeneratedAt: time.Now().UTC(),
119 - Self: self,
120 - Peers: nil,
121 - }
122 - if s.peerRegistry != nil {
123 - resp.Peers = s.peerRegistry.advertisedPeers()
124 - }
125 - if req.RootHost != "" && req.RootHost != s.rootHost {
126 - return resp, nil
127 - }
128 - if req.Name == "" {
129 - return resp, nil
130 - }
131 -
132 - hostname, err := utils.LeaseHostname(req.Name, s.rootHost)
133 - if err != nil {
134 - return types.DiscoverResponse{}, err
135 - }
136 -
137 - now := time.Now()
138 - var lease types.Lease
139 - ok := false
140 -
141 - s.registry.mu.RLock()
142 - if leaseID, found := s.registry.routes.LookupExact(hostname); found {
143 - record, found := s.registry.leaseByID[leaseID]
144 - if found && record != nil && !now.After(record.ExpiresAt) && !record.Metadata.Hide && s.registry.policy.IsLeaseRoutable(record.ID) {
145 - lease = record.Lease
146 - lease.Metadata = lease.Metadata.Copy()
147 - ok = true
148 - }
149 - }
150 - s.registry.mu.RUnlock()
151 -
152 - if !ok {
153 - return resp, nil
154 - }
155 -
156 - ownerAddress := lease.OwnerAddress
157 - if strings.TrimSpace(ownerAddress) == "" {
158 - ownerAddress = s.ownerIdentity.Address
159 - }
160 -
161 - resp.Service = &types.DiscoveredService{
162 - Found: true,
163 - Name: lease.Name,
164 - Hostname: lease.Hostname,
165 - ExpiresAt: lease.ExpiresAt,
166 - OwnerAddress: ownerAddress,
167 - RelayID: self.RelayID,
111 +func (s *Server) handleRelayDiscovery(w http.ResponseWriter, r *http.Request) {
112 + if r.Method != http.MethodGet {
113 + utils.WriteAPIError(w, http.StatusMethodNotAllowed, types.APIErrorCodeMethodNotAllowed, "method not allowed")
114 + return
115 }
169 - return resp, nil
170 -}
116
172 -func (s *Server) discoverySelfDescriptor() (types.RelayDescriptor, error) {
117 now := time.Now().UTC()
118 ingressAddr := s.rootHost
119 if s.cfg.SNIPort != 0 && s.cfg.SNIPort != 443 {
@@ -180,7 +124,7 @@ func (s *Server) discoverySelfDescriptor() (types.RelayDescriptor, error) {
124 strings.TrimSpace(s.wgConfig.Endpoint) != "" &&
125 strings.TrimSpace(s.wgConfig.OverlayIPv4) != ""
126
183 - descriptor := types.RelayDescriptor{
127 + self, err := discovery.SignedDescriptor(types.RelayDescriptor{
128 RelayID: s.cfg.PortalURL,
129 OwnerAddress: s.ownerIdentity.Address,
130 SignerPublicKey: s.ownerIdentity.PublicKey,
@@ -196,14 +140,26 @@ func (s *Server) discoverySelfDescriptor() (types.RelayDescriptor, error) {
140 SupportsWitness: false,
141 SupportsVPNExit: false,
142 StatusState: "healthy",
143 + WireGuardPublicKey: strings.TrimSpace(s.wgConfig.PublicKey),
144 + WireGuardEndpoint: strings.TrimSpace(s.wgConfig.Endpoint),
145 + OverlayIPv4: strings.TrimSpace(s.wgConfig.OverlayIPv4),
146 + OverlayCIDRs: append([]string(nil), s.wgConfig.OverlayCIDRs...),
147 + }, s.ownerIdentity.PrivateKey)
148 + if err != nil {
149 + utils.WriteAPIError(w, http.StatusInternalServerError, types.APIErrorCodeInternal, err.Error())
150 + return
151 + }
152 +
153 + resp := types.DiscoveryResponse{
154 + ProtocolVersion: 1,
155 + GeneratedAt: now,
156 + Self: self,
157 + Relays: nil,
158 }
200 - if supportsOverlayPeer {
201 - descriptor.WireGuardPublicKey = strings.TrimSpace(s.wgConfig.PublicKey)
202 - descriptor.WireGuardEndpoint = strings.TrimSpace(s.wgConfig.Endpoint)
203 - descriptor.OverlayIPv4 = strings.TrimSpace(s.wgConfig.OverlayIPv4)
204 - descriptor.OverlayCIDRs = append([]string(nil), s.wgConfig.OverlayCIDRs...)
159 + if s.relaySet != nil {
160 + resp.Relays = s.relaySet.AdvertisedDescriptors()
161 }
206 - return discovery.SignedDescriptor(descriptor, s.ownerIdentity.PrivateKey)
162 + utils.WriteAPIData(w, http.StatusOK, resp)
163 }
164
165 func (s *Server) handleDomain(w http.ResponseWriter, r *http.Request) {
portal/discovery/discovery.go
+123 -66
@@ -16,10 +16,14 @@ import (
16 "github.com/gosuda/portal/v2/utils"
17 )
18
19 -type Resolver func(context.Context, types.DiscoverRequest) (types.DiscoverResponse, error)
20 -
19 const defaultRequestTimeout = 15 * time.Second
20
21 +type RelayIdentity struct {
22 + RelayID string
23 + APIHTTPSAddr string
24 + SignerPublicKey string
25 +}
26 +
27 func NormalizeDescriptor(desc types.RelayDescriptor) (types.RelayDescriptor, error) {
28 desc.RelayID = strings.TrimSpace(desc.RelayID)
29 desc.SignerPublicKey = strings.ToLower(strings.TrimSpace(desc.SignerPublicKey))
@@ -158,59 +162,149 @@ func ValidateDescriptor(desc types.RelayDescriptor, now time.Time) (types.RelayD
162 return normalized, nil
163 }
164
161 -func ValidateResponse(resp types.DiscoverResponse, now time.Time) (types.RelayDescriptor, []types.RelayDescriptor, error) {
165 +func ValidateRelayDiscoveryResponse(resp types.DiscoveryResponse, now time.Time) (types.RelayDescriptor, []types.RelayDescriptor, error) {
166 self, err := ValidateDescriptor(resp.Self, now)
167 if err != nil {
168 return types.RelayDescriptor{}, nil, err
169 }
170
171 seen := map[string]struct{}{self.RelayID: {}}
168 - peers := make([]types.RelayDescriptor, 0, len(resp.Peers))
172 + relays := make([]types.RelayDescriptor, 0, len(resp.Relays))
173 var validateErr error
170 - for _, descriptor := range resp.Peers {
174 + for _, descriptor := range resp.Relays {
175 verified, err := ValidateDescriptor(descriptor, now)
176 if err != nil {
173 - validateErr = errors.Join(validateErr, fmt.Errorf("validate peer %q: %w", descriptor.RelayID, err))
177 + if validateErr == nil {
178 + validateErr = err
179 + }
180 continue
181 }
182 if _, ok := seen[verified.RelayID]; ok {
183 continue
184 }
185 seen[verified.RelayID] = struct{}{}
180 - peers = append(peers, verified)
186 + relays = append(relays, verified)
187 + }
188 + return self, relays, validateErr
189 +}
190 +
191 +func RelayIdentityFromDescriptor(desc types.RelayDescriptor) (RelayIdentity, error) {
192 + normalized, err := NormalizeDescriptor(desc)
193 + if err != nil {
194 + return RelayIdentity{}, err
195 + }
196 +
197 + identity := RelayIdentity{
198 + RelayID: strings.TrimSpace(normalized.RelayID),
199 + APIHTTPSAddr: strings.TrimSpace(normalized.APIHTTPSAddr),
200 + SignerPublicKey: strings.ToLower(strings.TrimSpace(normalized.SignerPublicKey)),
201 + }
202 + switch {
203 + case identity.RelayID == "":
204 + return RelayIdentity{}, errors.New("descriptor relay_id is required")
205 + case identity.APIHTTPSAddr == "":
206 + return RelayIdentity{}, errors.New("descriptor api_https_addr is required")
207 + case identity.SignerPublicKey == "":
208 + return RelayIdentity{}, errors.New("descriptor signer_public_key is required")
209 + }
210 + return identity, nil
211 +}
212 +
213 +func MatchTargetRelayIdentity(identity RelayIdentity, targetRelayID, targetURL string) error {
214 + targetRelayID = strings.TrimSpace(targetRelayID)
215 + if targetRelayID != "" && identity.RelayID != targetRelayID {
216 + return errors.New("descriptor relay_id does not match target relay")
217 + }
218 +
219 + targetURL = strings.TrimSpace(targetURL)
220 + if targetURL == "" {
221 + return nil
222 + }
223 +
224 + normalizedTargetURL, err := utils.NormalizeRelayURL(targetURL)
225 + if err != nil {
226 + return err
227 + }
228 + if identity.APIHTTPSAddr != normalizedTargetURL {
229 + return errors.New("descriptor api_https_addr does not match target url")
230 + }
231 + return nil
232 +}
233 +
234 +func MatchPinnedRelayIdentity(identity, pinned RelayIdentity) error {
235 + if relayID := strings.TrimSpace(pinned.RelayID); relayID != "" && identity.RelayID != relayID {
236 + return errors.New("descriptor relay_id does not match pinned relay id")
237 + }
238 + if apiURL := strings.TrimSpace(pinned.APIHTTPSAddr); apiURL != "" && identity.APIHTTPSAddr != apiURL {
239 + return errors.New("descriptor api_https_addr does not match pinned relay url")
240 + }
241 + if signerPublicKey := strings.ToLower(strings.TrimSpace(pinned.SignerPublicKey)); signerPublicKey != "" && identity.SignerPublicKey != signerPublicKey {
242 + return errors.New("descriptor signer_public_key does not match pinned signer")
243 + }
244 + return nil
245 +}
246 +
247 +func DiscoverRelayDiscovery(ctx context.Context, baseURL string, rootCAPEM []byte, httpClient *http.Client) (types.DiscoveryResponse, error) {
248 + resp, err := doGET[types.DiscoveryResponse](ctx, baseURL, types.PathDiscovery, nil, rootCAPEM, httpClient)
249 + if err != nil {
250 + return types.DiscoveryResponse{}, err
251 + }
252 + return resp, nil
253 +}
254 +
255 +func SeedDescriptor(apiURL string) (types.RelayDescriptor, error) {
256 + normalized, err := utils.NormalizeRelayURL(apiURL)
257 + if err != nil {
258 + return types.RelayDescriptor{}, err
259 + }
260 + return types.RelayDescriptor{
261 + RelayID: normalized,
262 + APIHTTPSAddr: normalized,
263 + Version: 1,
264 + }, nil
265 +}
266 +
267 +func RequireOverlayRelayDescriptor(desc types.RelayDescriptor) error {
268 + if !desc.SupportsOverlayPeer {
269 + return errors.New("descriptor does not support overlay peer")
270 + }
271 + if strings.TrimSpace(desc.WireGuardPublicKey) == "" {
272 + return errors.New("descriptor wireguard public key is required")
273 }
182 - return self, peers, validateErr
274 + if strings.TrimSpace(desc.WireGuardEndpoint) == "" {
275 + return errors.New("descriptor wireguard endpoint is required")
276 + }
277 + if strings.TrimSpace(desc.OverlayIPv4) == "" {
278 + return errors.New("descriptor overlay ipv4 is required")
279 + }
280 + return nil
281 }
282
185 -func Discover(ctx context.Context, baseURL string, req types.DiscoverRequest, rootCAPEM []byte, httpClient *http.Client) (types.DiscoverResponse, error) {
283 +func doGET[T any](ctx context.Context, baseURL, path string, query url.Values, rootCAPEM []byte, httpClient *http.Client) (T, error) {
284 + var zero T
285 +
286 baseURL = strings.TrimSpace(baseURL)
287 if baseURL == "" {
188 - return types.DiscoverResponse{}, errors.New("discovery base url is required")
288 + return zero, errors.New("discovery base url is required")
289 }
190 -
290 parsedBaseURL, err := url.Parse(baseURL)
291 if err != nil {
193 - return types.DiscoverResponse{}, fmt.Errorf("parse discovery base url: %w", err)
292 + return zero, fmt.Errorf("parse discovery base url: %w", err)
293 }
294 if parsedBaseURL.Host == "" {
196 - return types.DiscoverResponse{}, errors.New("discovery base url host is required")
295 + return zero, errors.New("discovery base url host is required")
296 }
297
199 - discoverURL := parsedBaseURL.ResolveReference(&url.URL{Path: types.PathDiscovery})
200 - query := discoverURL.Query()
201 - if req.RootHost != "" {
202 - query.Set("root_host", req.RootHost)
203 - }
204 - if req.Name != "" {
205 - query.Set("name", req.Name)
298 + requestURL := parsedBaseURL.ResolveReference(&url.URL{Path: path})
299 + if query != nil {
300 + requestURL.RawQuery = query.Encode()
301 }
207 - discoverURL.RawQuery = query.Encode()
302
303 client := httpClient
304 if client == nil {
305 rootCAs, err := keyless.RelayRootCAs(ctx, baseURL, parsedBaseURL.Hostname(), rootCAPEM)
306 if err != nil {
213 - return types.DiscoverResponse{}, err
307 + return zero, err
308 }
309 client = &http.Client{
310 Transport: &http.Transport{
@@ -230,64 +324,27 @@ func Discover(ctx context.Context, baseURL string, req types.DiscoverRequest, ro
324 client = &clone
325 }
326
233 - httpReq, err := http.NewRequestWithContext(ctx, http.MethodGet, discoverURL.String(), nil)
327 + httpReq, err := http.NewRequestWithContext(ctx, http.MethodGet, requestURL.String(), nil)
328 if err != nil {
235 - return types.DiscoverResponse{}, err
329 + return zero, err
330 }
331
332 resp, err := client.Do(httpReq)
333 if err != nil {
240 - return types.DiscoverResponse{}, err
334 + return zero, err
335 }
336 defer resp.Body.Close()
337
338 if resp.StatusCode < http.StatusOK || resp.StatusCode >= http.StatusMultipleChoices {
245 - return types.DiscoverResponse{}, utils.DecodeAPIRequestError(resp)
339 + return zero, utils.DecodeAPIRequestError(resp)
340 }
341
248 - envelope, err := utils.DecodeAPIEnvelope[types.DiscoverResponse](resp.Body)
342 + envelope, err := utils.DecodeAPIEnvelope[T](resp.Body)
343 if err != nil {
250 - return types.DiscoverResponse{}, fmt.Errorf("decode response: %w", err)
344 + return zero, fmt.Errorf("decode response: %w", err)
345 }
346 if !envelope.OK {
253 - return types.DiscoverResponse{}, utils.NewAPIRequestError(resp.StatusCode, envelope.Error)
347 + return zero, utils.NewAPIRequestError(resp.StatusCode, envelope.Error)
348 }
349 return envelope.Data, nil
350 }
257 -
258 -func ServeHTTP(w http.ResponseWriter, r *http.Request, resolver Resolver) {
259 - if r.Method != http.MethodGet {
260 - utils.WriteAPIError(w, http.StatusMethodNotAllowed, types.APIErrorCodeMethodNotAllowed, "method not allowed")
261 - return
262 - }
263 -
264 - req := types.DiscoverRequest{
265 - RootHost: r.URL.Query().Get("root_host"),
266 - Name: r.URL.Query().Get("name"),
267 - }
268 - req.RootHost = utils.NormalizeHostname(req.RootHost)
269 - req.Name = strings.TrimSpace(req.Name)
270 - if req.Name != "" {
271 - if req.RootHost == "" {
272 - utils.WriteAPIError(w, http.StatusBadRequest, types.APIErrorCodeInvalidRequest, "root host is required when name is set")
273 - return
274 - }
275 - name, err := utils.NormalizeDNSLabel(req.Name)
276 - if err != nil {
277 - utils.WriteAPIError(w, http.StatusBadRequest, types.APIErrorCodeInvalidRequest, err.Error())
278 - return
279 - }
280 - req.Name = name
281 - }
282 - if resolver == nil {
283 - utils.WriteAPIError(w, http.StatusInternalServerError, types.APIErrorCodeInternal, "discovery resolver is not configured")
284 - return
285 - }
286 -
287 - resp, err := resolver(r.Context(), req)
288 - if err != nil {
289 - utils.WriteAPIError(w, http.StatusInternalServerError, types.APIErrorCodeInternal, err.Error())
290 - return
291 - }
292 - utils.WriteAPIData(w, http.StatusOK, resp)
293 -}
portal/discovery/relayset.go new
+747
@@ -0,0 +1,747 @@
1 +package discovery
2 +
3 +import (
4 + "errors"
5 + "net/http"
6 + "reflect"
7 + "sort"
8 + "strings"
9 + "sync"
10 + "time"
11 +
12 + "github.com/rs/zerolog/log"
13 +
14 + "github.com/gosuda/portal/v2/types"
15 + "github.com/gosuda/portal/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 +
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 +
47 +// RelaySet owns the shared relay discovery view: known relay URLs, pinned relay
48 +// identities, 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 {
54 + mu sync.RWMutex
55 + knownRelayURLs []string
56 + pinnedByRelayID map[string]RelayIdentity
57 + relayIDsByURL map[string]string
58 + relays map[string]RelayView
59 + localByURL map[string]RelayLocalState
60 + lastStatusReachable map[string]bool
61 + lastStatusSummary RelaySummary
62 + haveLastStatus bool
63 +}
64 +
65 +func NewRelaySet() *RelaySet {
66 + return &RelaySet{
67 + pinnedByRelayID: make(map[string]RelayIdentity),
68 + relayIDsByURL: make(map[string]string),
69 + relays: make(map[string]RelayView),
70 + localByURL: make(map[string]RelayLocalState),
71 + }
72 +}
73 +
74 +func (s *RelaySet) trackedRelayURLs() []string {
75 + if s == nil {
76 + return nil
77 + }
78 +
79 + urls := make([]string, 0, len(s.knownRelayURLs)+len(s.relays))
80 + seen := make(map[string]struct{}, len(s.knownRelayURLs)+len(s.relays))
81 + for _, relayURL := range s.knownRelayURLs {
82 + relayURL = strings.TrimSpace(relayURL)
83 + if relayURL == "" {
84 + continue
85 + }
86 + if _, ok := seen[relayURL]; ok {
87 + continue
88 + }
89 + seen[relayURL] = struct{}{}
90 + urls = append(urls, relayURL)
91 + }
92 + for _, view := range s.relays {
93 + relayURL := strings.TrimSpace(view.Descriptor.APIHTTPSAddr)
94 + if relayURL == "" {
95 + continue
96 + }
97 + if _, ok := seen[relayURL]; ok {
98 + continue
99 + }
100 + seen[relayURL] = struct{}{}
101 + urls = append(urls, relayURL)
102 + }
103 + return urls
104 +}
105 +
106 +func (s *RelaySet) ActiveRelayURLs() []string {
107 + if s == nil {
108 + return nil
109 + }
110 + s.mu.RLock()
111 + defer s.mu.RUnlock()
112 + if len(s.knownRelayURLs) == 0 {
113 + return nil
114 + }
115 +
116 + out := make([]string, 0, len(s.knownRelayURLs))
117 + for _, relayURL := range s.knownRelayURLs {
118 + if state, ok := s.localByURL[relayURL]; ok && state.Banned {
119 + continue
120 + }
121 + out = append(out, relayURL)
122 + }
123 + if len(out) == 0 {
124 + return nil
125 + }
126 + return out
127 +}
128 +
129 +func (s *RelaySet) logStatusChange() {
130 + var currentReachable map[string]bool
131 + trackedRelayURLs := s.trackedRelayURLs()
132 + if len(trackedRelayURLs) > 0 {
133 + currentReachable = make(map[string]bool, len(trackedRelayURLs))
134 + for _, relayURL := range trackedRelayURLs {
135 + state := s.localByURL[relayURL]
136 + currentReachable[relayURL] = !state.Banned && state.Reachable
137 + }
138 + }
139 + summary := RelaySummary{}
140 + for _, relayURL := range trackedRelayURLs {
141 + summary.Known++
142 + state := s.localByURL[relayURL]
143 + if state.Banned {
144 + summary.Banned++
145 + continue
146 + }
147 + if state.Bootstrap {
148 + summary.Bootstrap++
149 + }
150 + if state.Advertised {
151 + summary.Advertised++
152 + }
153 + if state.Expired {
154 + summary.Expired++
155 + }
156 + if state.Reachable {
157 + summary.Reachable++
158 + } else {
159 + summary.Unreachable++
160 + }
161 + relayID := s.relayIDsByURL[relayURL]
162 + view, ok := s.relays[relayID]
163 + if ok && !state.Bootstrap && !state.Expired && view.Descriptor.SupportsOverlayPeer {
164 + summary.Syncable++
165 + }
166 + }
167 + if s.haveLastStatus && summary == s.lastStatusSummary && reflect.DeepEqual(currentReachable, s.lastStatusReachable) {
168 + return
169 + }
170 +
171 + activated := make([]string, 0)
172 + deactivated := make([]string, 0)
173 + for relayURL, reachable := range currentReachable {
174 + if s.lastStatusReachable == nil || s.lastStatusReachable[relayURL] == reachable {
175 + continue
176 + }
177 + if reachable {
178 + activated = append(activated, relayURL)
179 + } else {
180 + deactivated = append(deactivated, relayURL)
181 + }
182 + }
183 + for relayURL, reachable := range s.lastStatusReachable {
184 + if _, ok := currentReachable[relayURL]; ok || !reachable {
185 + continue
186 + }
187 + deactivated = append(deactivated, relayURL)
188 + }
189 +
190 + event := log.Info().
191 + Int("banned", summary.Banned).
192 + Int("bootstrap", summary.Bootstrap).
193 + Int("advertised", summary.Advertised).
194 + Int("expired", summary.Expired).
195 + Int("syncable", summary.Syncable).
196 + Int("reachable", summary.Reachable).
197 + Int("unreachable", summary.Unreachable)
198 + if len(activated) > 0 {
199 + event = event.Strs("activated", activated)
200 + }
201 + if len(deactivated) > 0 {
202 + event = event.Strs("deactivated", deactivated)
203 + }
204 + event.Msg("relay status")
205 + s.lastStatusReachable = currentReachable
206 + s.lastStatusSummary = summary
207 + s.haveLastStatus = true
208 +}
209 +
210 +func (s *RelaySet) BootstrapDescriptors() []types.RelayDescriptor {
211 + if s == nil {
212 + return nil
213 + }
214 + s.mu.RLock()
215 + defer s.mu.RUnlock()
216 + if len(s.knownRelayURLs) == 0 {
217 + return nil
218 + }
219 +
220 + out := make([]types.RelayDescriptor, 0, len(s.knownRelayURLs))
221 + for _, relayURL := range s.knownRelayURLs {
222 + state, ok := s.localByURL[relayURL]
223 + if !ok || !state.Bootstrap {
224 + continue
225 + }
226 + if relayID, ok := s.relayIDsByURL[relayURL]; ok {
227 + if view, ok := s.relays[relayID]; ok && strings.TrimSpace(view.Descriptor.APIHTTPSAddr) != "" {
228 + out = append(out, view.Descriptor)
229 + continue
230 + }
231 + }
232 + out = append(out, types.RelayDescriptor{
233 + RelayID: relayURL,
234 + APIHTTPSAddr: relayURL,
235 + Version: 1,
236 + })
237 + }
238 + if len(out) == 0 {
239 + return nil
240 + }
241 + return out
242 +}
243 +
244 +func (s *RelaySet) BanRelayURL(relayURL, reason string) bool {
245 + if s == nil {
246 + return false
247 + }
248 + s.mu.Lock()
249 + defer s.mu.Unlock()
250 + relayURL = strings.TrimSpace(relayURL)
251 + if relayURL == "" {
252 + return false
253 + }
254 +
255 + state := s.localByURL[relayURL]
256 + changed := !state.Banned || strings.TrimSpace(state.BanReason) != strings.TrimSpace(reason)
257 + state.Banned = true
258 + state.BanReason = strings.TrimSpace(reason)
259 + state.Reachable = false
260 + s.localByURL[relayURL] = state
261 + if changed {
262 + s.logStatusChange()
263 + }
264 + return changed
265 +}
266 +
267 +func (s *RelaySet) MarkRelayUnreachable(relayURL string) bool {
268 + if s == nil {
269 + return false
270 + }
271 + s.mu.Lock()
272 + defer s.mu.Unlock()
273 + relayURL = strings.TrimSpace(relayURL)
274 + if relayURL == "" {
275 + return false
276 + }
277 +
278 + state := s.localByURL[relayURL]
279 + if state.Banned {
280 + return false
281 + }
282 + if !state.Reachable {
283 + return false
284 + }
285 + state.Reachable = false
286 + s.localByURL[relayURL] = state
287 + s.logStatusChange()
288 + return true
289 +}
290 +
291 +func (s *RelaySet) MarkRelayReachable(relayURL string, now time.Time) bool {
292 + if s == nil {
293 + return false
294 + }
295 + s.mu.Lock()
296 + defer s.mu.Unlock()
297 + relayURL = strings.TrimSpace(relayURL)
298 + if relayURL == "" {
299 + return false
300 + }
301 + if now.IsZero() {
302 + now = time.Now().UTC()
303 + }
304 +
305 + state := s.localByURL[relayURL]
306 + changed := !state.Reachable || state.ConsecutiveFailures != 0 || state.LastSuccessAt != now
307 + state.Reachable = true
308 + state.ConsecutiveFailures = 0
309 + state.LastSuccessAt = now
310 + s.localByURL[relayURL] = state
311 + if changed {
312 + s.logStatusChange()
313 + }
314 + return changed
315 +}
316 +
317 +func (s *RelaySet) MarkRelayFailure(relayURL string, now time.Time) RelayLocalState {
318 + if s == nil {
319 + return RelayLocalState{}
320 + }
321 + s.mu.Lock()
322 + defer s.mu.Unlock()
323 + relayURL = strings.TrimSpace(relayURL)
324 + if relayURL == "" {
325 + return RelayLocalState{}
326 + }
327 + if now.IsZero() {
328 + now = time.Now().UTC()
329 + }
330 +
331 + state := s.localByURL[relayURL]
332 + state.Reachable = false
333 + state.ConsecutiveFailures++
334 + state.LastFailureAt = now
335 + s.localByURL[relayURL] = state
336 + s.logStatusChange()
337 + return state
338 +}
339 +
340 +func (s *RelaySet) AdvertisedDescriptors() []types.RelayDescriptor {
341 + if s == nil {
342 + return nil
343 + }
344 + s.mu.RLock()
345 + defer s.mu.RUnlock()
346 + if len(s.relays) == 0 {
347 + return nil
348 + }
349 +
350 + out := make([]types.RelayDescriptor, 0, len(s.relays))
351 + for _, view := range s.relays {
352 + state := s.localByURL[view.Descriptor.APIHTTPSAddr]
353 + if !state.Advertised || state.Expired || strings.TrimSpace(view.Descriptor.APIHTTPSAddr) == "" {
354 + continue
355 + }
356 + out = append(out, view.Descriptor)
357 + }
358 + if len(out) == 0 {
359 + return nil
360 + }
361 + sort.Slice(out, func(i, j int) bool {
362 + return out[i].APIHTTPSAddr < out[j].APIHTTPSAddr
363 + })
364 + return out
365 +}
366 +
367 +func (s *RelaySet) SyncableDescriptors() []types.RelayDescriptor {
368 + if s == nil {
369 + return nil
370 + }
371 + s.mu.RLock()
372 + defer s.mu.RUnlock()
373 + if len(s.relays) == 0 {
374 + return nil
375 + }
376 +
377 + out := make([]types.RelayDescriptor, 0, len(s.relays))
378 + for _, view := range s.relays {
379 + state := s.localByURL[view.Descriptor.APIHTTPSAddr]
380 + if state.Bootstrap || state.Expired || !view.Descriptor.SupportsOverlayPeer {
381 + continue
382 + }
383 + out = append(out, view.Descriptor)
384 + }
385 + if len(out) == 0 {
386 + return nil
387 + }
388 + sort.Slice(out, func(i, j int) bool {
389 + return out[i].APIHTTPSAddr < out[j].APIHTTPSAddr
390 + })
391 + return out
392 +}
393 +
394 +func (s *RelaySet) Snapshot() map[string]types.RelayState {
395 + if s == nil {
396 + return nil
397 + }
398 + s.mu.RLock()
399 + defer s.mu.RUnlock()
400 + if len(s.relays) == 0 {
401 + return nil
402 + }
403 +
404 + snapshot := make(map[string]types.RelayState, len(s.relays))
405 + for relayID, view := range s.relays {
406 + localState := s.localByURL[view.Descriptor.APIHTTPSAddr]
407 + snapshot[relayID] = types.RelayState{
408 + Descriptor: view.Descriptor,
409 + Bootstrap: localState.Bootstrap,
410 + Advertised: localState.Advertised,
411 + Expired: localState.Expired,
412 + FirstSeenAt: view.FirstSeenAt,
413 + LastSeenAt: view.LastSeenAt,
414 + ConsecutiveFailures: localState.ConsecutiveFailures,
415 + }
416 + }
417 + return snapshot
418 +}
419 +
420 +func (s *RelaySet) ReplaceKnownRelayURLs(relayURLs []string) {
421 + if s == nil {
422 + return
423 + }
424 + s.mu.Lock()
425 + defer s.mu.Unlock()
426 + filtered := make([]string, 0, len(relayURLs))
427 + for _, relayURL := range relayURLs {
428 + relayURL = strings.TrimSpace(relayURL)
429 + if relayURL == "" {
430 + continue
431 + }
432 + duplicate := false
433 + for _, existing := range filtered {
434 + if existing == relayURL {
435 + duplicate = true
436 + break
437 + }
438 + }
439 + if duplicate {
440 + continue
441 + }
442 + filtered = append(filtered, relayURL)
443 + }
444 + s.knownRelayURLs = append([]string(nil), filtered...)
445 +}
446 +
447 +func (s *RelaySet) pinTarget(targetRelayID, targetURL string, desc types.RelayDescriptor) error {
448 + if s == nil {
449 + return nil
450 + }
451 + identity, err := RelayIdentityFromDescriptor(desc)
452 + if err != nil {
453 + return err
454 + }
455 + if err := MatchTargetRelayIdentity(identity, targetRelayID, targetURL); err != nil {
456 + return err
457 + }
458 + if err := s.matchPinned(identity); err != nil {
459 + return err
460 + }
461 + s.pin(identity)
462 + return nil
463 +}
464 +
465 +func (s *RelaySet) registerDescriptor(desc types.RelayDescriptor, now time.Time) (string, bool, bool, error) {
466 + if s == nil {
467 + return "", false, false, nil
468 + }
469 + identity, err := RelayIdentityFromDescriptor(desc)
470 + if err != nil {
471 + return "", false, false, err
472 + }
473 + if err := s.matchPinned(identity); err != nil {
474 + return "", false, false, err
475 + }
476 + s.pin(identity)
477 +
478 + if now.IsZero() {
479 + now = time.Now().UTC()
480 + }
481 +
482 + view, ok := s.relays[identity.RelayID]
483 + added := !ok
484 + if !ok {
485 + view.FirstSeenAt = now
486 + }
487 + previousDescriptor := view.Descriptor
488 + view.Descriptor = desc
489 + view.LastSeenAt = now
490 + s.relays[identity.RelayID] = view
491 +
492 + changed := added || !reflect.DeepEqual(previousDescriptor, desc)
493 + return identity.RelayID, added, changed, nil
494 +}
495 +
496 +func relayDiscoveryURLs(selfDescriptor types.RelayDescriptor, relayDescriptors []types.RelayDescriptor) []string {
497 + relayURLs := make([]string, 0, 1+len(relayDescriptors))
498 + if apiURL := strings.TrimSpace(selfDescriptor.APIHTTPSAddr); apiURL != "" {
499 + relayURLs = append(relayURLs, apiURL)
500 + }
501 + for _, relayDescriptor := range relayDescriptors {
502 + if apiURL := strings.TrimSpace(relayDescriptor.APIHTTPSAddr); apiURL != "" {
503 + relayURLs = append(relayURLs, apiURL)
504 + }
505 + }
506 + if len(relayURLs) == 0 {
507 + return nil
508 + }
509 + return relayURLs
510 +}
511 +
512 +func (s *RelaySet) applyDiscoveryDescriptors(targetRelayID, targetURL string, selfDescriptor types.RelayDescriptor, relayDescriptors []types.RelayDescriptor, now time.Time) (relaySetChanged bool, addedRelayCount int, err error) {
513 + if s == nil {
514 + return false, 0, nil
515 + }
516 + if strings.TrimSpace(targetRelayID) == "" {
517 + return false, 0, errors.New("target relay id is required")
518 + }
519 + if now.IsZero() {
520 + now = time.Now().UTC()
521 + }
522 + if err := s.pinTarget(targetRelayID, targetURL, selfDescriptor); err != nil {
523 + return false, 0, err
524 + }
525 +
526 + apply := func(desc types.RelayDescriptor, advertise, countAdded bool) error {
527 + _, added, descriptorChanged, err := s.registerDescriptor(desc, now)
528 + if err != nil {
529 + return err
530 + }
531 + localState := s.localByURL[desc.APIHTTPSAddr]
532 + wasAdvertised := localState.Advertised
533 + wasExpired := localState.Expired
534 + if advertise {
535 + localState.Advertised = true
536 + }
537 + localState.Expired = false
538 + s.localByURL[desc.APIHTTPSAddr] = localState
539 +
540 + changed := added || descriptorChanged || advertise && !wasAdvertised || wasExpired
541 + if added && countAdded {
542 + addedRelayCount++
543 + }
544 + if changed {
545 + relaySetChanged = true
546 + }
547 + return nil
548 + }
549 +
550 + if err := apply(selfDescriptor, true, false); err != nil {
551 + return false, 0, err
552 + }
553 + for _, relayDescriptor := range relayDescriptors {
554 + if err := apply(relayDescriptor, false, true); err != nil {
555 + return false, 0, err
556 + }
557 + }
558 + state := s.localByURL[selfDescriptor.APIHTTPSAddr]
559 + state.Reachable = true
560 + state.ConsecutiveFailures = 0
561 + state.LastSuccessAt = now
562 + s.localByURL[selfDescriptor.APIHTTPSAddr] = state
563 + s.logStatusChange()
564 + return relaySetChanged, addedRelayCount, nil
565 +}
566 +
567 +func (s *RelaySet) ApplyRelayDiscoveryResponse(targetRelayID, targetURL string, resp types.DiscoveryResponse, now time.Time) (relayURLs []string, relaySetChanged bool, addedRelayCount int, warnErr error, err error) {
568 + selfDescriptor, relayDescriptors, validateErr := ValidateRelayDiscoveryResponse(resp, now)
569 + warnErr = validateErr
570 + if selfDescriptor.RelayID == "" {
571 + return nil, false, 0, warnErr, validateErr
572 + }
573 + if selfDescriptor.RelayID != strings.TrimSpace(targetRelayID) {
574 + return nil, false, 0, warnErr, errors.New("relay discovery response relay_id mismatch")
575 + }
576 + s.mu.Lock()
577 + relaySetChanged, addedRelayCount, err = s.applyDiscoveryDescriptors(targetRelayID, targetURL, selfDescriptor, relayDescriptors, now)
578 + s.mu.Unlock()
579 + if err != nil {
580 + return nil, false, 0, warnErr, err
581 + }
582 + return relayDiscoveryURLs(selfDescriptor, relayDescriptors), relaySetChanged, addedRelayCount, warnErr, nil
583 +}
584 +
585 +func (s *RelaySet) ApplyOverlayRelayDiscoveryResponse(targetRelayID, targetURL string, resp types.DiscoveryResponse, now time.Time) (relayURLs []string, relaySetChanged bool, addedRelayCount int, warnErr error, err error) {
586 + selfDescriptor, relayDescriptors, validateErr := ValidateRelayDiscoveryResponse(resp, now)
587 + warnErr = validateErr
588 + if selfDescriptor.RelayID == "" {
589 + return nil, false, 0, warnErr, validateErr
590 + }
591 + if selfDescriptor.RelayID != strings.TrimSpace(targetRelayID) {
592 + return nil, false, 0, warnErr, errors.New("relay discovery response relay_id mismatch")
593 + }
594 + if err := RequireOverlayRelayDescriptor(selfDescriptor); err != nil {
595 + return nil, false, 0, warnErr, err
596 + }
597 +
598 + filteredRelayDescriptors := make([]types.RelayDescriptor, 0, len(relayDescriptors))
599 + for _, relayDescriptor := range relayDescriptors {
600 + if err := RequireOverlayRelayDescriptor(relayDescriptor); err != nil {
601 + if warnErr == nil {
602 + warnErr = err
603 + }
604 + continue
605 + }
606 + filteredRelayDescriptors = append(filteredRelayDescriptors, relayDescriptor)
607 + }
608 +
609 + s.mu.Lock()
610 + relaySetChanged, addedRelayCount, err = s.applyDiscoveryDescriptors(targetRelayID, targetURL, selfDescriptor, filteredRelayDescriptors, now)
611 + s.mu.Unlock()
612 + if err != nil {
613 + return nil, false, 0, warnErr, err
614 + }
615 + return relayDiscoveryURLs(selfDescriptor, filteredRelayDescriptors), relaySetChanged, addedRelayCount, warnErr, nil
616 +}
617 +
618 +func (s *RelaySet) RegisterBootstrapRelayURLs(inputs []string, now time.Time) ([]string, error) {
619 + if s == nil || len(inputs) == 0 {
620 + return nil, nil
621 + }
622 +
623 + normalized, err := utils.NormalizeRelayURLs(inputs...)
624 + if err != nil {
625 + return nil, err
626 + }
627 + normalized, err = utils.ExcludeLocalRelayURLs(normalized...)
628 + if err != nil {
629 + return nil, err
630 + }
631 + if len(normalized) == 0 {
632 + return nil, nil
633 + }
634 + if now.IsZero() {
635 + now = time.Now().UTC()
636 + }
637 +
638 + s.mu.Lock()
639 + defer s.mu.Unlock()
640 +
641 + existing := make(map[string]struct{}, len(s.knownRelayURLs))
642 + for _, relayURL := range s.knownRelayURLs {
643 + existing[relayURL] = struct{}{}
644 + }
645 + added := make([]string, 0, len(normalized))
646 + for _, relayURL := range normalized {
647 + if _, ok := existing[relayURL]; ok {
648 + continue
649 + }
650 + existing[relayURL] = struct{}{}
651 + s.knownRelayURLs = append(s.knownRelayURLs, relayURL)
652 + added = append(added, relayURL)
653 + }
654 + for _, relayURL := range normalized {
655 + state := s.localByURL[relayURL]
656 + state.Bootstrap = true
657 + state.Reachable = false
658 + s.localByURL[relayURL] = state
659 + if descriptor, err := SeedDescriptor(relayURL); err == nil {
660 + _, _, _, _ = s.registerDescriptor(descriptor, now)
661 + }
662 + }
663 + s.logStatusChange()
664 + if len(added) == 0 {
665 + return nil, nil
666 + }
667 + return added, nil
668 +}
669 +
670 +func (s *RelaySet) RecordDiscoveryFailure(relayID, relayURL string, err error, recoveryFailures int, now time.Time) (expired bool, expireReason string, consecutiveFailures int) {
671 + if s == nil {
672 + return false, "", 0
673 + }
674 + relayID = strings.TrimSpace(relayID)
675 + if relayID == "" {
676 + return false, "", 0
677 + }
678 + relayURL = strings.TrimSpace(relayURL)
679 + if relayURL == "" {
680 + return false, "", 0
681 + }
682 + if now.IsZero() {
683 + now = time.Now().UTC()
684 + }
685 +
686 + s.mu.Lock()
687 + defer s.mu.Unlock()
688 +
689 + view, ok := s.relays[relayID]
690 + if !ok {
691 + return false, "", 0
692 + }
693 +
694 + localState := s.localByURL[relayURL]
695 + localState.Reachable = false
696 + localState.ConsecutiveFailures++
697 + localState.LastFailureAt = now
698 + s.localByURL[relayURL] = localState
699 + s.logStatusChange()
700 + if !localState.Expired && localState.ConsecutiveFailures >= recoveryFailures {
701 + state := s.localByURL[view.Descriptor.APIHTTPSAddr]
702 + state.Expired = true
703 + s.localByURL[view.Descriptor.APIHTTPSAddr] = state
704 + s.logStatusChange()
705 + return true, "recovery", localState.ConsecutiveFailures
706 + }
707 +
708 + var apiErr *types.APIRequestError
709 + if errors.As(err, &apiErr) &&
710 + (apiErr.StatusCode == http.StatusForbidden ||
711 + apiErr.StatusCode == http.StatusNotFound ||
712 + apiErr.StatusCode == http.StatusGone) {
713 + state := s.localByURL[view.Descriptor.APIHTTPSAddr]
714 + state.Expired = true
715 + s.localByURL[view.Descriptor.APIHTTPSAddr] = state
716 + s.logStatusChange()
717 + return true, "status", localState.ConsecutiveFailures
718 + }
719 + return false, "", localState.ConsecutiveFailures
720 +}
721 +
722 +func (s *RelaySet) matchPinned(identity RelayIdentity) error {
723 + if s == nil {
724 + return nil
725 + }
726 + if pinned, ok := s.pinnedByRelayID[identity.RelayID]; ok {
727 + if err := MatchPinnedRelayIdentity(identity, pinned); err != nil {
728 + return err
729 + }
730 + }
731 + if pinnedRelayID, ok := s.relayIDsByURL[identity.APIHTTPSAddr]; ok && pinnedRelayID != identity.RelayID {
732 + return MatchPinnedRelayIdentity(identity, RelayIdentity{
733 + RelayID: pinnedRelayID,
734 + APIHTTPSAddr: identity.APIHTTPSAddr,
735 + SignerPublicKey: "",
736 + })
737 + }
738 + return nil
739 +}
740 +
741 +func (s *RelaySet) pin(identity RelayIdentity) {
742 + if s == nil {
743 + return
744 + }
745 + s.pinnedByRelayID[identity.RelayID] = identity
746 + s.relayIDsByURL[identity.APIHTTPSAddr] = identity.RelayID
747 +}
portal/lease.go
+1 -4
@@ -18,7 +18,6 @@ type leaseRegistry struct {
18 routes *routeTable
19 leaseByID map[string]*leaseRecord
20 policy *policy.Runtime
21 - onExpired func(*leaseRecord) // called for each expired lease during cleanup
21 mu sync.RWMutex
22 }
23
@@ -193,9 +192,7 @@ func (r *leaseRegistry) Touch(leaseID, clientIP string, now time.Time) *leaseRec
192
193 func (r *leaseRegistry) cleanupExpired(now time.Time) {
194 for _, lease := range r.removeExpired(now) {
196 - if r.onExpired != nil {
197 - r.onExpired(lease)
198 - }
195 + lease.Close()
196 }
197 }
198
portal/lease_test.go
-3
@@ -153,9 +153,6 @@ func TestLeaseRegistryCleanupExpiredClosesBroker(t *testing.T) {
153 t.Parallel()
154
155 registry := newLeaseRegistry(policy.NewRuntime())
156 - registry.onExpired = func(r *leaseRecord) {
157 - r.Close()
158 - }
156 record := &leaseRecord{
157 Lease: types.Lease{
158 ID: "lease_expired",
portal/peer.go deleted
-407
@@ -1,407 +0,0 @@
1 -package portal
2 -
3 -import (
4 - "errors"
5 - "net/http"
6 - "reflect"
7 - "sort"
8 - "strings"
9 - "sync"
10 - "time"
11 -
12 - "github.com/gosuda/portal/v2/types"
13 - "github.com/gosuda/portal/v2/utils"
14 -)
15 -
16 -type peerRecord struct {
17 - bootstrap bool
18 - seedURL string
19 - pinnedSignerPublicKey string
20 - state types.PeerState
21 -}
22 -
23 -// peerRegistry keeps only verified relay descriptors and bootstrap hints.
24 -// It is not a source of truth; trust comes from seed URLs plus signer pinning.
25 -type peerRegistry struct {
26 - mu sync.RWMutex
27 - peers map[string]peerRecord
28 -}
29 -
30 -type peerRegistrationResult struct {
31 - updated bool
32 - addedHintCount int
33 - peerSetChanged bool
34 -}
35 -
36 -type peerFailureOutcome struct {
37 - expired bool
38 - expireReason string
39 - consecutiveFailures int
40 -}
41 -
42 -func newPeerRegistry() *peerRegistry {
43 - return &peerRegistry{
44 - peers: make(map[string]peerRecord),
45 - }
46 -}
47 -
48 -func (r *peerRegistry) lookup(relayID string) (types.PeerState, bool, bool) {
49 - if strings.TrimSpace(relayID) == "" {
50 - return types.PeerState{}, false, false
51 - }
52 -
53 - r.mu.RLock()
54 - defer r.mu.RUnlock()
55 -
56 - record, ok := r.peers[relayID]
57 - if !ok {
58 - return types.PeerState{}, false, false
59 - }
60 - return record.state, strings.TrimSpace(record.pinnedSignerPublicKey) != "", true
61 -}
62 -
63 -func seedDescriptor(apiURL string) (types.RelayDescriptor, error) {
64 - normalized, err := utils.NormalizeRelayURL(apiURL)
65 - if err != nil {
66 - return types.RelayDescriptor{}, err
67 - }
68 - return types.RelayDescriptor{
69 - RelayID: normalized,
70 - APIHTTPSAddr: normalized,
71 - Version: 1,
72 - }, nil
73 -}
74 -
75 -func requireOverlayPeerDescriptor(desc types.RelayDescriptor) error {
76 - if !desc.SupportsOverlayPeer {
77 - return errors.New("descriptor does not support overlay peer")
78 - }
79 - if strings.TrimSpace(desc.WireGuardPublicKey) == "" {
80 - return errors.New("descriptor wireguard public key is required")
81 - }
82 - if strings.TrimSpace(desc.WireGuardEndpoint) == "" {
83 - return errors.New("descriptor wireguard endpoint is required")
84 - }
85 - if strings.TrimSpace(desc.OverlayIPv4) == "" {
86 - return errors.New("descriptor overlay ipv4 is required")
87 - }
88 - return nil
89 -}
90 -
91 -func (r *peerRegistry) registerBootstrapURLs(inputs []string) ([]string, error) {
92 - if len(inputs) == 0 {
93 - return nil, nil
94 - }
95 -
96 - normalized, err := utils.NormalizeRelayURLs(inputs...)
97 - if err != nil {
98 - return nil, err
99 - }
100 - normalized, err = utils.ExcludeLocalRelayURLs(normalized...)
101 - if err != nil {
102 - return nil, err
103 - }
104 -
105 - now := time.Now().UTC()
106 - added := make([]string, 0, len(normalized))
107 -
108 - r.mu.Lock()
109 - defer r.mu.Unlock()
110 -
111 - for _, apiURL := range normalized {
112 - descriptor, err := seedDescriptor(apiURL)
113 - if err != nil {
114 - return nil, err
115 - }
116 -
117 - record, ok := r.peers[descriptor.RelayID]
118 - if !ok {
119 - r.peers[descriptor.RelayID] = peerRecord{
120 - bootstrap: true,
121 - seedURL: descriptor.APIHTTPSAddr,
122 - state: types.PeerState{
123 - Descriptor: descriptor,
124 - State: types.PeerStateKnown,
125 - FirstSeenAt: now,
126 - LastSeenAt: now,
127 - },
128 - }
129 - added = append(added, descriptor.APIHTTPSAddr)
130 - continue
131 - }
132 -
133 - record.bootstrap = true
134 - if strings.TrimSpace(record.seedURL) == "" {
135 - record.seedURL = descriptor.APIHTTPSAddr
136 - }
137 - if strings.TrimSpace(record.state.Descriptor.APIHTTPSAddr) == "" {
138 - record.state.Descriptor.APIHTTPSAddr = descriptor.APIHTTPSAddr
139 - }
140 - if strings.TrimSpace(record.state.Descriptor.RelayID) == "" {
141 - record.state.Descriptor.RelayID = descriptor.RelayID
142 - }
143 - record.state.LastSeenAt = now
144 - r.peers[descriptor.RelayID] = record
145 - }
146 -
147 - return added, nil
148 -}
149 -
150 -func (r *peerRegistry) snapshot() map[string]types.PeerState {
151 - r.mu.RLock()
152 - defer r.mu.RUnlock()
153 -
154 - out := make(map[string]types.PeerState, len(r.peers))
155 - for relayID, record := range r.peers {
156 - out[relayID] = record.state
157 - }
158 - return out
159 -}
160 -
161 -func (r *peerRegistry) bootstrapPeers() []types.RelayDescriptor {
162 - r.mu.RLock()
163 - defer r.mu.RUnlock()
164 -
165 - out := make([]types.RelayDescriptor, 0, len(r.peers))
166 - for _, record := range r.peers {
167 - if !record.bootstrap || strings.TrimSpace(record.state.Descriptor.APIHTTPSAddr) == "" {
168 - continue
169 - }
170 - out = append(out, record.state.Descriptor)
171 - }
172 - sort.Slice(out, func(i, j int) bool {
173 - return out[i].APIHTTPSAddr < out[j].APIHTTPSAddr
174 - })
175 - return out
176 -}
177 -
178 -func (r *peerRegistry) advertisedPeers() []types.RelayDescriptor {
179 - r.mu.RLock()
180 - defer r.mu.RUnlock()
181 -
182 - out := make([]types.RelayDescriptor, 0, len(r.peers))
183 - for _, record := range r.peers {
184 - if record.state.State != types.PeerStateAdvertised {
185 - continue
186 - }
187 - out = append(out, record.state.Descriptor)
188 - }
189 - sort.Slice(out, func(i, j int) bool {
190 - return out[i].APIHTTPSAddr < out[j].APIHTTPSAddr
191 - })
192 - return out
193 -}
194 -
195 -func (r *peerRegistry) syncablePeers() []types.RelayDescriptor {
196 - r.mu.RLock()
197 - defer r.mu.RUnlock()
198 -
199 - out := make([]types.RelayDescriptor, 0, len(r.peers))
200 - for _, record := range r.peers {
201 - if record.bootstrap ||
202 - record.state.State == types.PeerStateExpired ||
203 - !record.state.Descriptor.SupportsOverlayPeer {
204 - continue
205 - }
206 - out = append(out, record.state.Descriptor)
207 - }
208 - sort.Slice(out, func(i, j int) bool {
209 - return out[i].APIHTTPSAddr < out[j].APIHTTPSAddr
210 - })
211 - return out
212 -}
213 -
214 -func (r *peerRegistry) pin(relayID, seedURL string, desc types.RelayDescriptor) error {
215 - relayID = strings.TrimSpace(relayID)
216 - if relayID == "" {
217 - return errors.New("relay id is required")
218 - }
219 - normalizedSeedURL, err := utils.NormalizeRelayURL(seedURL)
220 - if err != nil {
221 - return err
222 - }
223 - if strings.TrimSpace(desc.APIHTTPSAddr) == "" {
224 - return errors.New("descriptor api_https_addr is required")
225 - }
226 - if desc.APIHTTPSAddr != normalizedSeedURL {
227 - return errors.New("descriptor api_https_addr does not match seed url")
228 - }
229 - if strings.TrimSpace(desc.SignerPublicKey) == "" {
230 - return errors.New("descriptor signer_public_key is required")
231 - }
232 -
233 - now := time.Now().UTC()
234 -
235 - r.mu.Lock()
236 - defer r.mu.Unlock()
237 -
238 - record, ok := r.peers[relayID]
239 - if !ok {
240 - record = peerRecord{
241 - state: types.PeerState{
242 - FirstSeenAt: now,
243 - },
244 - }
245 - }
246 - if record.seedURL != "" && record.seedURL != normalizedSeedURL {
247 - return errors.New("seed url does not match cached relay url")
248 - }
249 - if record.pinnedSignerPublicKey != "" && record.pinnedSignerPublicKey != desc.SignerPublicKey {
250 - return errors.New("descriptor signer_public_key does not match pinned signer")
251 - }
252 -
253 - record.seedURL = normalizedSeedURL
254 - record.pinnedSignerPublicKey = desc.SignerPublicKey
255 - r.peers[relayID] = record
256 - return nil
257 -}
258 -
259 -func (r *peerRegistry) register(desc types.RelayDescriptor, advertise bool) (bool, bool, error) {
260 - if strings.TrimSpace(desc.RelayID) == "" {
261 - return false, false, errors.New("relay id is required")
262 - }
263 -
264 - now := time.Now().UTC()
265 -
266 - r.mu.Lock()
267 - defer r.mu.Unlock()
268 -
269 - record, ok := r.peers[desc.RelayID]
270 - added := !ok
271 - if !ok {
272 - record = peerRecord{
273 - state: types.PeerState{
274 - FirstSeenAt: now,
275 - },
276 - }
277 - }
278 -
279 - if strings.TrimSpace(record.seedURL) == "" {
280 - record.seedURL = strings.TrimSpace(desc.APIHTTPSAddr)
281 - } else if strings.TrimSpace(desc.APIHTTPSAddr) != "" && record.seedURL != strings.TrimSpace(desc.APIHTTPSAddr) {
282 - return false, false, errors.New("descriptor api_https_addr does not match cached seed url")
283 - }
284 - if record.pinnedSignerPublicKey != "" && record.pinnedSignerPublicKey != desc.SignerPublicKey {
285 - return false, false, errors.New("descriptor signer_public_key does not match pinned signer")
286 - }
287 -
288 - previousState := record.state.State
289 - previousDescriptor := record.state.Descriptor
290 - switch {
291 - case advertise:
292 - record.state.State = types.PeerStateAdvertised
293 - case record.state.State == types.PeerStateAdvertised:
294 - record.state.State = types.PeerStateAdvertised
295 - default:
296 - record.state.State = types.PeerStateVerified
297 - }
298 -
299 - record.state.Descriptor = desc
300 - record.state.LastSeenAt = now
301 - record.state.ConsecutiveFailures = 0
302 - r.peers[desc.RelayID] = record
303 -
304 - changed := added ||
305 - previousState != record.state.State ||
306 - !reflect.DeepEqual(previousDescriptor, record.state.Descriptor)
307 - return added, changed, nil
308 -}
309 -
310 -func (r *peerRegistry) registerDiscoveredPeers(targetRelayID, targetURL string, selfDescriptor types.RelayDescriptor, peerDescriptors []types.RelayDescriptor) (peerRegistrationResult, error) {
311 - if strings.TrimSpace(targetRelayID) == "" {
312 - return peerRegistrationResult{}, errors.New("target relay id is required")
313 - }
314 - if err := r.pin(targetRelayID, targetURL, selfDescriptor); err != nil {
315 - return peerRegistrationResult{}, err
316 - }
317 -
318 - added, changed, err := r.register(selfDescriptor, true)
319 - if err != nil {
320 - return peerRegistrationResult{}, err
321 - }
322 -
323 - result := peerRegistrationResult{
324 - updated: added || changed,
325 - peerSetChanged: changed,
326 - }
327 -
328 - for _, peerDescriptor := range peerDescriptors {
329 - hintAdded, hintChanged, err := r.register(peerDescriptor, false)
330 - if err != nil {
331 - return peerRegistrationResult{}, err
332 - }
333 - result.peerSetChanged = result.peerSetChanged || hintChanged
334 - if hintAdded || hintChanged {
335 - result.updated = true
336 - result.addedHintCount++
337 - }
338 - }
339 -
340 - return result, nil
341 -}
342 -
343 -func (r *peerRegistry) fail(relayID string) {
344 - if strings.TrimSpace(relayID) == "" {
345 - return
346 - }
347 -
348 - r.mu.Lock()
349 - defer r.mu.Unlock()
350 -
351 - record, ok := r.peers[relayID]
352 - if !ok {
353 - return
354 - }
355 - record.state.ConsecutiveFailures++
356 - r.peers[relayID] = record
357 -}
358 -
359 -func (r *peerRegistry) expire(relayID string) bool {
360 - if strings.TrimSpace(relayID) == "" {
361 - return false
362 - }
363 -
364 - r.mu.Lock()
365 - defer r.mu.Unlock()
366 -
367 - record, ok := r.peers[relayID]
368 - if !ok {
369 - return false
370 - }
371 - if record.state.State == types.PeerStateExpired {
372 - return false
373 - }
374 - record.state.State = types.PeerStateExpired
375 - r.peers[relayID] = record
376 - return true
377 -}
378 -
379 -func (r *peerRegistry) recordFailure(relayID string, err error, recoveryFailures int) peerFailureOutcome {
380 - r.fail(relayID)
381 -
382 - result := peerFailureOutcome{}
383 - state, _, ok := r.lookup(relayID)
384 - if ok &&
385 - state.State != types.PeerStateExpired &&
386 - state.ConsecutiveFailures >= recoveryFailures {
387 - if removed := r.expire(relayID); removed {
388 - result.expired = true
389 - result.expireReason = "recovery"
390 - result.consecutiveFailures = state.ConsecutiveFailures
391 - return result
392 - }
393 - }
394 -
395 - var apiErr *types.APIRequestError
396 - if errors.As(err, &apiErr) &&
397 - (apiErr.StatusCode == http.StatusForbidden ||
398 - apiErr.StatusCode == http.StatusNotFound ||
399 - apiErr.StatusCode == http.StatusGone) {
400 - if removed := r.expire(relayID); removed {
401 - result.expired = true
402 - result.expireReason = "status"
403 - }
404 - }
405 -
406 - return result
407 -}
portal/server.go
+127 -180
@@ -83,7 +83,7 @@ type Server struct {
83 cfg ServerConfig
84 rootHost string
85 trustedProxyCIDRs []*net.IPNet
86 - peerRegistry *peerRegistry
86 + relaySet *discovery.RelaySet
87 shutdownOnce sync.Once
88 }
89
@@ -112,6 +112,14 @@ func NewServer(cfg ServerConfig) (*Server, error) {
112 return nil, fmt.Errorf("normalize bootstraps: %w", err)
113 }
114 cfg.Bootstraps = bootstraps
115 + generatedWireGuardPrivateKey := ""
116 + if cfg.DiscoveryEnabled && strings.TrimSpace(cfg.WireGuardPrivateKey) == "" {
117 + generatedWireGuardPrivateKey, err = utils.GenerateWireGuardPrivateKey()
118 + if err != nil {
119 + return nil, err
120 + }
121 + cfg.WireGuardPrivateKey = generatedWireGuardPrivateKey
122 + }
123 wgConfig, err := wireguard.NormalizeConfig(rootHost, wireguard.Config{
124 PrivateKey: cfg.WireGuardPrivateKey,
125 PublicKey: cfg.WireGuardPublicKey,
@@ -123,8 +131,11 @@ func NewServer(cfg ServerConfig) (*Server, error) {
131 if err != nil {
132 return nil, err
133 }
126 - if cfg.DiscoveryEnabled && strings.TrimSpace(wgConfig.PrivateKey) == "" {
127 - return nil, errors.New("wireguard private key is required when discovery is enabled")
134 + if generatedWireGuardPrivateKey != "" {
135 + log.Warn().
136 + Str("wireguard_public_key", wgConfig.PublicKey).
137 + Str("wireguard_private_key", generatedWireGuardPrivateKey).
138 + Msg("generated wireguard private key; set WIREGUARD_PRIVATE_KEY to preserve relay identity")
139 }
140
141 portMin, portMax := 0, 0
@@ -164,16 +175,10 @@ func NewServer(cfg ServerConfig) (*Server, error) {
175 trustedProxyCIDRs: trustedProxyCIDRs,
176 }
177
167 - // Tear down all lease resources when leases expire via TTL janitor.
168 - registry.onExpired = func(record *leaseRecord) {
169 - if record != nil {
170 - record.Close()
171 - }
172 - }
173 -
178 if cfg.DiscoveryEnabled {
175 - s.peerRegistry = newPeerRegistry()
176 - if _, err := s.peerRegistry.registerBootstrapURLs(cfg.Bootstraps); err != nil {
179 + s.relaySet = discovery.NewRelaySet()
180 + _, err = s.relaySet.RegisterBootstrapRelayURLs(cfg.Bootstraps, time.Now().UTC())
181 + if err != nil {
182 return nil, err
183 }
184 }
@@ -226,9 +231,9 @@ func (s *Server) Start(ctx context.Context, apiMux *http.ServeMux) error {
231 s.group = group
232
233 if s.wgConfig.PrivateKey != "" {
229 - var snapshot map[string]types.PeerState
230 - if s.peerRegistry != nil {
231 - snapshot = s.peerRegistry.snapshot()
234 + var snapshot map[string]types.RelayState
235 + if s.relaySet != nil {
236 + snapshot = s.relaySet.Snapshot()
237 }
238 peerMux := http.NewServeMux()
239 peerMux.HandleFunc(types.PathRoot, s.handleRoot)
@@ -238,7 +243,7 @@ func (s *Server) Start(ctx context.Context, apiMux *http.ServeMux) error {
243 http.NotFound(w, r)
244 return
245 }
241 - discovery.ServeHTTP(w, r, s.discover)
246 + s.handleRelayDiscovery(w, r)
247 })
248 overlay, err := wireguard.NewOverlay(s.wgConfig, peerMux)
249 if err != nil {
@@ -268,14 +273,8 @@ func (s *Server) Start(ctx context.Context, apiMux *http.ServeMux) error {
273 group.Go(func() error { return s.runSNIListener(groupCtx) })
274 group.Go(func() error { return s.registry.RunJanitor(groupCtx, 5*time.Second) })
275 if s.DiscoveryEnabled() {
271 - group.Go(func() error { return s.runDiscoveryLoop(groupCtx) })
276 + group.Go(func() error { return s.runRelayDiscoveryLoop(groupCtx) })
277 }
273 - group.Go(func() error {
274 - <-groupCtx.Done()
275 - shutdownCtx, cancel := context.WithTimeout(context.Background(), 5*time.Second)
276 - defer cancel()
277 - return s.Shutdown(shutdownCtx)
278 - })
278 s.acmeManager.Start(serverCtx)
279
280 if s.cfg.UDPPortCount > 0 {
@@ -283,6 +282,12 @@ func (s *Server) Start(ctx context.Context, apiMux *http.ServeMux) error {
282 log.Warn().Err(err).Msg("quic tunnel listener disabled")
283 }
284 }
285 + group.Go(func() error {
286 + <-groupCtx.Done()
287 + shutdownCtx, cancel := context.WithTimeout(context.Background(), 5*time.Second)
288 + defer cancel()
289 + return s.Shutdown(shutdownCtx)
290 + })
291
292 return nil
293 }
@@ -559,183 +564,125 @@ func (s *Server) runQUICTunnelListener(listener *quic.Listener) error {
564 }
565 }
566
562 -func (s *Server) applyDiscoveryResponse(targetRelayID, targetURL string, resp types.DiscoverResponse, requireSelfOverlay bool) (updated bool, addedHintCount int, warnErr error, err error) {
563 - if strings.TrimSpace(targetRelayID) == "" {
564 - return false, 0, nil, errors.New("target relay id is required")
565 - }
566 -
567 - now := time.Now().UTC()
568 - selfDescriptor, peerDescriptors, warnErr := discovery.ValidateResponse(resp, now)
569 - if selfDescriptor.RelayID == "" {
570 - return false, 0, nil, errors.Join(warnErr, errors.New("discover response is missing self descriptor"))
571 - }
572 - if selfDescriptor.RelayID != targetRelayID {
573 - return false, 0, nil, errors.Join(warnErr, errors.New("discover response relay_id mismatch"))
574 - }
575 - if requireSelfOverlay {
576 - if err := requireOverlayPeerDescriptor(selfDescriptor); err != nil {
577 - return false, 0, nil, errors.Join(warnErr, err)
578 - }
579 - }
580 -
581 - filteredPeerDescriptors := make([]types.RelayDescriptor, 0, len(peerDescriptors))
582 -
583 - for _, peerDescriptor := range peerDescriptors {
584 - if err := requireOverlayPeerDescriptor(peerDescriptor); err != nil {
585 - warnErr = errors.Join(warnErr, fmt.Errorf("record hint %q: %w", peerDescriptor.RelayID, err))
586 - continue
587 - }
588 - filteredPeerDescriptors = append(filteredPeerDescriptors, peerDescriptor)
589 - }
590 -
591 - result, err := s.peerRegistry.registerDiscoveredPeers(targetRelayID, targetURL, selfDescriptor, filteredPeerDescriptors)
592 - if err != nil {
593 - return false, 0, nil, errors.Join(warnErr, err)
594 - }
595 - updated = result.updated
596 - addedHintCount = result.addedHintCount
597 -
598 - if result.peerSetChanged && s.overlay != nil {
599 - if err := s.overlay.Sync(s.cfg.PortalURL, s.peerRegistry.snapshot()); err != nil {
600 - warnErr = errors.Join(warnErr, err)
601 - }
602 - }
567 +func (s *Server) runRelayDiscoveryLoop(ctx context.Context) error {
568 + ticker := time.NewTicker(defaultDiscoveryInterval)
569 + defer ticker.Stop()
570
604 - return updated, addedHintCount, warnErr, nil
605 -}
571 + for {
572 + bootstraps := s.relaySet.BootstrapDescriptors()
573
607 -func (s *Server) runBootstrapDiscoveryPass(ctx context.Context) {
608 - for _, bootstrap := range s.peerRegistry.bootstrapPeers() {
609 - resp, err := discovery.Discover(ctx, bootstrap.APIHTTPSAddr, types.DiscoverRequest{}, nil, nil)
610 - if err != nil {
611 - if ctx.Err() != nil {
612 - return
574 + for _, bootstrap := range bootstraps {
575 + resp, err := discovery.DiscoverRelayDiscovery(ctx, bootstrap.APIHTTPSAddr, nil, nil)
576 + if err != nil {
577 + if ctx.Err() != nil {
578 + return nil
579 + }
580 + s.relaySet.MarkRelayFailure(bootstrap.APIHTTPSAddr, time.Now().UTC())
581 + log.Warn().
582 + Err(err).
583 + Str("relay", bootstrap.APIHTTPSAddr).
584 + Msg("bootstrap relay discovery failed")
585 + continue
586 }
614 - s.peerRegistry.fail(bootstrap.RelayID)
615 - log.Warn().
616 - Err(err).
617 - Str("peer", bootstrap.APIHTTPSAddr).
618 - Msg("bootstrap discovery failed")
619 - continue
620 - }
621 -
622 - updated, addedHintCount, warnErr, err := s.applyDiscoveryResponse(bootstrap.RelayID, bootstrap.APIHTTPSAddr, resp, false)
623 - if err != nil {
624 - s.peerRegistry.fail(bootstrap.RelayID)
625 - log.Warn().
626 - Err(err).
627 - Str("peer", bootstrap.APIHTTPSAddr).
628 - Msg("bootstrap discovery failed")
629 - continue
630 - }
587
632 - if updated || warnErr != nil {
633 - event := log.Info()
634 - if warnErr != nil {
635 - event = log.Warn().Err(warnErr)
588 + now := time.Now().UTC()
589 + var relaySetChanged bool
590 + var warnErr error
591 + _, relaySetChanged, _, warnErr, err = s.relaySet.ApplyRelayDiscoveryResponse(bootstrap.RelayID, bootstrap.APIHTTPSAddr, resp, now)
592 + if relaySetChanged && s.overlay != nil {
593 + if syncErr := s.overlay.Sync(s.cfg.PortalURL, s.relaySet.Snapshot()); syncErr != nil {
594 + if warnErr == nil {
595 + warnErr = syncErr
596 + }
597 + }
598 }
637 - event = event.
638 - Str("peer", bootstrap.APIHTTPSAddr).
639 - Int("bootstrap_count", len(s.peerRegistry.bootstrapPeers())).
640 - Int("peer_count", len(s.peerRegistry.syncablePeers())).
641 - Int("advertised_count", len(s.peerRegistry.advertisedPeers()))
642 - if addedHintCount > 0 {
643 - event = event.Int("added_hint_count", addedHintCount)
599 + if err != nil {
600 + s.relaySet.MarkRelayFailure(bootstrap.APIHTTPSAddr, time.Now().UTC())
601 + log.Warn().
602 + Err(err).
603 + Str("relay", bootstrap.APIHTTPSAddr).
604 + Msg("bootstrap relay discovery failed")
605 + continue
606 }
607 +
608 if warnErr != nil {
646 - event.Msg("bootstrap discovery completed with warnings")
647 - } else {
648 - event.Msg("bootstrap discovery updated")
609 + log.Warn().
610 + Err(warnErr).
611 + Str("relay", bootstrap.APIHTTPSAddr).
612 + Msg("bootstrap relay discovery completed with warnings")
613 }
614 }
651 - }
652 -}
615 + if ctx.Err() != nil {
616 + return nil
617 + }
618
654 -func (s *Server) runOverlayPeerDiscoveryPass(ctx context.Context) {
655 - if s.overlay == nil {
656 - return
657 - }
658 - overlayClient := s.overlay.Client()
659 - if overlayClient == nil {
660 - return
661 - }
619 + if s.overlay != nil {
620 + overlayClient := s.overlay.Client()
621 + syncableRelays := s.relaySet.SyncableDescriptors()
622
663 - for _, peer := range s.peerRegistry.syncablePeers() {
664 - var failureErr error
623 + for _, relay := range syncableRelays {
624 + var failureErr error
625
666 - if err := requireOverlayPeerDescriptor(peer); err != nil {
667 - failureErr = err
668 - } else {
669 - discoverURL := "http://" + net.JoinHostPort(peer.OverlayIPv4, fmt.Sprintf("%d", wireguard.DefaultPeerAPIHTTPPort))
670 - resp, err := discovery.Discover(ctx, discoverURL, types.DiscoverRequest{}, nil, overlayClient)
671 - if err != nil {
672 - if ctx.Err() != nil {
673 - return
674 - }
675 - failureErr = err
676 - } else {
677 - updated, addedHintCount, warnErr, err := s.applyDiscoveryResponse(peer.RelayID, peer.APIHTTPSAddr, resp, true)
678 - if err != nil {
626 + if err := discovery.RequireOverlayRelayDescriptor(relay); err != nil {
627 failureErr = err
628 } else {
681 - if warnErr != nil {
682 - log.Warn().
683 - Err(warnErr).
684 - Str("peer", peer.APIHTTPSAddr).
685 - Int("bootstrap_count", len(s.peerRegistry.bootstrapPeers())).
686 - Int("peer_count", len(s.peerRegistry.syncablePeers())).
687 - Int("advertised_count", len(s.peerRegistry.advertisedPeers())).
688 - Int("added_hint_count", addedHintCount).
689 - Msg("overlay peer discovery completed with warnings")
690 - continue
691 - }
692 -
693 - if updated {
694 - event := log.Info().
695 - Str("peer", peer.APIHTTPSAddr).
696 - Int("bootstrap_count", len(s.peerRegistry.bootstrapPeers())).
697 - Int("peer_count", len(s.peerRegistry.syncablePeers())).
698 - Int("advertised_count", len(s.peerRegistry.advertisedPeers()))
699 - if addedHintCount > 0 {
700 - event = event.Int("added_hint_count", addedHintCount)
629 + discoverURL := "http://" + net.JoinHostPort(relay.OverlayIPv4, fmt.Sprintf("%d", wireguard.DefaultPeerAPIHTTPPort))
630 + resp, err := discovery.DiscoverRelayDiscovery(ctx, discoverURL, nil, overlayClient)
631 + if err != nil {
632 + if ctx.Err() != nil {
633 + return nil
634 + }
635 + failureErr = err
636 + } else {
637 + now := time.Now().UTC()
638 + var relaySetChanged bool
639 + var warnErr error
640 + var snapshot map[string]types.RelayState
641 + _, relaySetChanged, _, warnErr, err = s.relaySet.ApplyOverlayRelayDiscoveryResponse(relay.RelayID, relay.APIHTTPSAddr, resp, now)
642 + if relaySetChanged {
643 + snapshot = s.relaySet.Snapshot()
644 + if syncErr := s.overlay.Sync(s.cfg.PortalURL, snapshot); syncErr != nil {
645 + if warnErr == nil {
646 + warnErr = syncErr
647 + }
648 + }
649 + }
650 + if err != nil {
651 + failureErr = err
652 + } else {
653 + if warnErr != nil {
654 + log.Warn().
655 + Err(warnErr).
656 + Str("relay", relay.APIHTTPSAddr).
657 + Msg("overlay relay discovery completed with warnings")
658 + continue
659 + }
660 +
661 + continue
662 }
702 - event.Msg("overlay peer discovery updated")
663 }
704 - continue
664 }
706 - }
707 - }
708 -
709 - result := s.peerRegistry.recordFailure(peer.RelayID, failureErr, defaultWGRecoveryFailures)
710 - if result.expired && s.overlay != nil {
711 - failureErr = errors.Join(failureErr, s.overlay.Sync(s.cfg.PortalURL, s.peerRegistry.snapshot()))
712 - }
665 + expired, expireReason, consecutiveFailures := s.relaySet.RecordDiscoveryFailure(relay.RelayID, relay.APIHTTPSAddr, failureErr, defaultWGRecoveryFailures, time.Now().UTC())
666 + if expired {
667 + if syncErr := s.overlay.Sync(s.cfg.PortalURL, s.relaySet.Snapshot()); syncErr != nil && failureErr == nil {
668 + failureErr = syncErr
669 + }
670 + }
671
714 - event := log.Warn().
715 - Err(failureErr).
716 - Str("peer", peer.APIHTTPSAddr)
717 - if result.expired {
718 - event = event.
719 - Bool("expired", true).
720 - Str("reason", result.expireReason)
721 - if result.consecutiveFailures > 0 {
722 - event = event.Int("consecutive_failures", result.consecutiveFailures)
672 + event := log.Warn().
673 + Err(failureErr).
674 + Str("relay", relay.APIHTTPSAddr)
675 + if expired {
676 + event = event.
677 + Bool("expired", true).
678 + Str("reason", expireReason)
679 + if consecutiveFailures > 0 {
680 + event = event.Int("consecutive_failures", consecutiveFailures)
681 + }
682 + }
683 + event.Msg("overlay relay discovery failed")
684 }
685 }
725 - event.Msg("overlay peer discovery failed")
726 - }
727 -}
728 -
729 -func (s *Server) runDiscoveryLoop(ctx context.Context) error {
730 - ticker := time.NewTicker(defaultDiscoveryInterval)
731 - defer ticker.Stop()
732 -
733 - for {
734 - s.runBootstrapDiscoveryPass(ctx)
735 - if ctx.Err() != nil {
736 - return nil
737 - }
738 - s.runOverlayPeerDiscoveryPass(ctx)
686 if ctx.Err() != nil {
687 return nil
688 }
portal/server_test.go
+117 -60
@@ -7,6 +7,7 @@ import (
7 "net"
8 "net/http"
9 "reflect"
10 + "sort"
11 "strings"
12 "testing"
13 "time"
@@ -60,18 +61,27 @@ func mustSignedRelayDescriptor(t *testing.T, ownerPrivateKey, relayURL string) t
61 return desc
62 }
63
63 -func TestNewServerRequiresWireGuardWhenDiscoveryEnabled(t *testing.T) {
64 +func TestNewServerGeneratesWireGuardWhenDiscoveryEnabled(t *testing.T) {
65 t.Parallel()
66
66 - _, err := NewServer(ServerConfig{
67 + server, err := NewServer(ServerConfig{
68 PortalURL: "https://portal.example.com",
69 DiscoveryEnabled: true,
70 })
70 - if err == nil {
71 - t.Fatal("NewServer() error = nil, want wireguard requirement error")
71 + if err != nil {
72 + t.Fatalf("NewServer() error = %v", err)
73 + }
74 + if server.wgConfig.PrivateKey == "" {
75 + t.Fatal("WireGuardPrivateKey = empty, want generated key")
76 }
73 - if !strings.Contains(err.Error(), "wireguard private key is required when discovery is enabled") {
74 - t.Fatalf("NewServer() error = %v, want wireguard requirement error", err)
77 + if server.wgConfig.PublicKey == "" {
78 + t.Fatal("WireGuardPublicKey = empty, want derived key")
79 + }
80 + if server.wgConfig.Endpoint == "" {
81 + t.Fatal("WireGuardEndpoint = empty, want derived endpoint")
82 + }
83 + if server.wgConfig.OverlayIPv4 == "" {
84 + t.Fatal("OverlayIPv4 = empty, want derived overlay address")
85 }
86 }
87
@@ -311,11 +321,11 @@ func TestServerUpsertDiscoverySeedURLsSkipsLocalRelayHosts(t *testing.T) {
321 t.Fatalf("NewServer() error = %v", err)
322 }
323
314 - added, err := server.peerRegistry.registerBootstrapURLs([]string{
324 + added, err := server.relaySet.RegisterBootstrapRelayURLs([]string{
325 "https://localhost:4017",
326 "https://relay-a.example.com",
327 "https://127.0.0.1:4017",
318 - })
328 + }, time.Now().UTC())
329 if err != nil {
330 t.Fatalf("UpsertSeedURLs() error = %v", err)
331 }
@@ -327,12 +337,17 @@ func TestServerUpsertDiscoverySeedURLsSkipsLocalRelayHosts(t *testing.T) {
337 if err != nil {
338 t.Fatalf("ExcludeLocalRelayURLs() error = %v", err)
339 }
330 - knownURLs := make([]string, 0, len(server.peerRegistry.bootstrapPeers()))
331 - for _, descriptor := range server.peerRegistry.bootstrapPeers() {
332 - if apiURL := strings.TrimSpace(descriptor.APIHTTPSAddr); apiURL != "" {
333 - knownURLs = append(knownURLs, apiURL)
340 + bootstrapDescriptors := server.relaySet.BootstrapDescriptors()
341 + syncableDescriptors := server.relaySet.SyncableDescriptors()
342 + advertisedDescriptors := server.relaySet.AdvertisedDescriptors()
343 + knownURLs := make([]string, 0, len(bootstrapDescriptors))
344 + for _, descriptor := range bootstrapDescriptors {
345 + if strings.TrimSpace(descriptor.APIHTTPSAddr) == "" {
346 + continue
347 }
348 + knownURLs = append(knownURLs, descriptor.APIHTTPSAddr)
349 }
350 + sort.Strings(knownURLs)
351 knownURLs, err = utils.ExcludeLocalRelayURLs(knownURLs...)
352 if err != nil {
353 t.Fatalf("ExcludeLocalRelayURLs() known error = %v", err)
@@ -340,11 +355,11 @@ func TestServerUpsertDiscoverySeedURLsSkipsLocalRelayHosts(t *testing.T) {
355 if !reflect.DeepEqual(knownURLs, knownRelayURLs) {
356 t.Fatalf("BootstrapDescriptors() = %v, want [%q %q]", knownURLs, "https://bootstrap.example.com", "https://relay-a.example.com")
357 }
343 - if len(server.peerRegistry.syncablePeers()) != 0 {
344 - t.Fatalf("syncablePeers() = %v, want empty before direct confirmation", server.peerRegistry.syncablePeers())
358 + if len(syncableDescriptors) != 0 {
359 + t.Fatalf("syncable count = %d, want 0 before direct confirmation", len(syncableDescriptors))
360 }
346 - if len(server.peerRegistry.advertisedPeers()) != 0 {
347 - t.Fatalf("advertisedPeers() = %v, want empty before direct confirmation", server.peerRegistry.advertisedPeers())
361 + if len(advertisedDescriptors) != 0 {
362 + t.Fatalf("advertised count = %d, want 0 before direct confirmation", len(advertisedDescriptors))
363 }
364 }
365
@@ -367,33 +382,66 @@ func TestServerRecordVerifiedDiscoveryPeerRequiresDirectConfirmation(t *testing.
382 bootstrapDesc := mustSignedRelayDescriptor(t, ownerPrivateKey, "https://bootstrap.example.com")
383 relayADesc := mustSignedRelayDescriptor(t, ownerPrivateKey, "https://relay-a.example.com")
384
370 - added, changed, err := server.peerRegistry.register(bootstrapDesc, true)
385 + applyDiscovery := func(targetRelayID, targetURL string, resp types.DiscoveryResponse, requireSelfOverlay bool) (bool, int, error, error) {
386 + now := time.Now().UTC()
387 + if requireSelfOverlay {
388 + _, updated, added, warnErr, err := server.relaySet.ApplyOverlayRelayDiscoveryResponse(targetRelayID, targetURL, resp, now)
389 + return updated, added, warnErr, err
390 + }
391 + _, updated, added, warnErr, err := server.relaySet.ApplyRelayDiscoveryResponse(targetRelayID, targetURL, resp, now)
392 + return updated, added, warnErr, err
393 + }
394 +
395 + resultUpdated, resultAdded, warnErr, err := applyDiscovery(
396 + bootstrapDesc.RelayID,
397 + bootstrapDesc.APIHTTPSAddr,
398 + types.DiscoveryResponse{Self: bootstrapDesc},
399 + false,
400 + )
401 if err != nil {
372 - t.Fatalf("RecordVerified() error = %v", err)
402 + t.Fatalf("applyRelayDiscoveryResponse() bootstrap error = %v", err)
403 + }
404 + if warnErr != nil {
405 + t.Fatalf("applyRelayDiscoveryResponse() bootstrap warn = %v, want nil", warnErr)
406 }
374 - if added {
375 - t.Fatal("RecordVerified() added = true, want false for seeded bootstrap")
407 + if resultAdded != 0 {
408 + t.Fatalf("applyRelayDiscoveryResponse() bootstrap added = %d, want 0 for seeded bootstrap", resultAdded)
409 }
377 - if !changed {
378 - t.Fatal("RecordVerified() changed = false, want true")
410 + if !resultUpdated {
411 + t.Fatal("applyRelayDiscoveryResponse() bootstrap updated = false, want true")
412 }
413
381 - added, changed, err = server.peerRegistry.register(relayADesc, false)
414 + resultUpdated, resultAdded, warnErr, err = applyDiscovery(
415 + bootstrapDesc.RelayID,
416 + bootstrapDesc.APIHTTPSAddr,
417 + types.DiscoveryResponse{Self: bootstrapDesc, Relays: []types.RelayDescriptor{relayADesc}},
418 + false,
419 + )
420 if err != nil {
383 - t.Fatalf("RecordVerified() hinted error = %v", err)
421 + t.Fatalf("applyRelayDiscoveryResponse() hinted error = %v", err)
422 + }
423 + if warnErr != nil {
424 + t.Fatalf("applyRelayDiscoveryResponse() hinted warn = %v, want nil", warnErr)
425 }
385 - if !added {
386 - t.Fatal("RecordVerified() hinted add = false, want true")
426 + if resultAdded != 1 {
427 + t.Fatalf("applyRelayDiscoveryResponse() hinted added = %d, want 1", resultAdded)
428 }
388 - if !changed {
389 - t.Fatal("RecordVerified() hinted changed = false, want true")
429 + if !resultUpdated {
430 + t.Fatal("applyRelayDiscoveryResponse() hinted updated = false, want true")
431 }
391 - knownURLs := make([]string, 0, len(server.peerRegistry.bootstrapPeers()))
392 - for _, descriptor := range server.peerRegistry.bootstrapPeers() {
393 - if apiURL := strings.TrimSpace(descriptor.APIHTTPSAddr); apiURL != "" {
394 - knownURLs = append(knownURLs, apiURL)
432 + snapshot := server.relaySet.Snapshot()
433 + if len(snapshot) != 2 {
434 + t.Fatalf("Snapshot() size = %d, want 2 after hinted relay registration", len(snapshot))
435 + }
436 + bootstrapDescriptors := server.relaySet.BootstrapDescriptors()
437 + knownURLs := make([]string, 0, len(bootstrapDescriptors))
438 + for _, descriptor := range bootstrapDescriptors {
439 + if strings.TrimSpace(descriptor.APIHTTPSAddr) == "" {
440 + continue
441 }
442 + knownURLs = append(knownURLs, descriptor.APIHTTPSAddr)
443 }
444 + sort.Strings(knownURLs)
445 knownURLs, err = utils.ExcludeLocalRelayURLs(knownURLs...)
446 if err != nil {
447 t.Fatalf("ExcludeLocalRelayURLs() known error = %v", err)
@@ -401,12 +449,15 @@ func TestServerRecordVerifiedDiscoveryPeerRequiresDirectConfirmation(t *testing.
449 if !reflect.DeepEqual(knownURLs, []string{"https://bootstrap.example.com"}) {
450 t.Fatalf("BootstrapDescriptors() = %v, want [%q]", knownURLs, "https://bootstrap.example.com")
451 }
404 - syncableURLs := make([]string, 0, len(server.peerRegistry.syncablePeers()))
405 - for _, descriptor := range server.peerRegistry.syncablePeers() {
406 - if apiURL := strings.TrimSpace(descriptor.APIHTTPSAddr); apiURL != "" {
407 - syncableURLs = append(syncableURLs, apiURL)
452 + syncableDescriptors := server.relaySet.SyncableDescriptors()
453 + syncableURLs := make([]string, 0, len(syncableDescriptors))
454 + for _, descriptor := range syncableDescriptors {
455 + if strings.TrimSpace(descriptor.APIHTTPSAddr) == "" {
456 + continue
457 }
458 + syncableURLs = append(syncableURLs, descriptor.APIHTTPSAddr)
459 }
460 + sort.Strings(syncableURLs)
461 syncableURLs, err = utils.ExcludeLocalRelayURLs(syncableURLs...)
462 if err != nil {
463 t.Fatalf("ExcludeLocalRelayURLs() syncable error = %v", err)
@@ -414,12 +465,15 @@ func TestServerRecordVerifiedDiscoveryPeerRequiresDirectConfirmation(t *testing.
465 if !reflect.DeepEqual(syncableURLs, []string{"https://relay-a.example.com"}) {
466 t.Fatalf("SyncablePeerDescriptors() = %v, want [%q]", syncableURLs, "https://relay-a.example.com")
467 }
417 - advertisedURLs := make([]string, 0, len(server.peerRegistry.advertisedPeers()))
418 - for _, descriptor := range server.peerRegistry.advertisedPeers() {
419 - if apiURL := strings.TrimSpace(descriptor.APIHTTPSAddr); apiURL != "" {
420 - advertisedURLs = append(advertisedURLs, apiURL)
468 + advertisedDescriptors := server.relaySet.AdvertisedDescriptors()
469 + advertisedURLs := make([]string, 0, len(advertisedDescriptors))
470 + for _, descriptor := range advertisedDescriptors {
471 + if strings.TrimSpace(descriptor.APIHTTPSAddr) == "" {
472 + continue
473 }
474 + advertisedURLs = append(advertisedURLs, descriptor.APIHTTPSAddr)
475 }
476 + sort.Strings(advertisedURLs)
477 advertisedURLs, err = utils.ExcludeLocalRelayURLs(advertisedURLs...)
478 if err != nil {
479 t.Fatalf("ExcludeLocalRelayURLs() advertised error = %v", err)
@@ -428,30 +482,33 @@ func TestServerRecordVerifiedDiscoveryPeerRequiresDirectConfirmation(t *testing.
482 t.Fatalf("AdvertisedDescriptors() = %v, want [%q]", advertisedURLs, "https://bootstrap.example.com")
483 }
484
431 - snapshot := server.peerRegistry.snapshot()
432 - if snapshot[bootstrapDesc.RelayID].State != types.PeerStateAdvertised {
433 - t.Fatalf("bootstrap state = %q, want %q", snapshot[bootstrapDesc.RelayID].State, types.PeerStateAdvertised)
434 - }
435 - if snapshot[relayADesc.RelayID].State != types.PeerStateVerified {
436 - t.Fatalf("relay-a state = %q, want %q", snapshot[relayADesc.RelayID].State, types.PeerStateVerified)
437 - }
438 -
439 - added, changed, err = server.peerRegistry.register(relayADesc, true)
485 + resultUpdated, resultAdded, warnErr, err = applyDiscovery(
486 + relayADesc.RelayID,
487 + relayADesc.APIHTTPSAddr,
488 + types.DiscoveryResponse{Self: relayADesc},
489 + true,
490 + )
491 if err != nil {
441 - t.Fatalf("RecordVerified() second error = %v", err)
492 + t.Fatalf("applyRelayDiscoveryResponse() confirm error = %v", err)
493 + }
494 + if warnErr != nil {
495 + t.Fatalf("applyRelayDiscoveryResponse() confirm warn = %v, want nil", warnErr)
496 }
443 - if added {
444 - t.Fatal("RecordVerified() second add = true, want false")
497 + if resultAdded != 0 {
498 + t.Fatalf("applyRelayDiscoveryResponse() confirm added = %d, want 0", resultAdded)
499 }
446 - if !changed {
447 - t.Fatal("RecordVerified() second changed = false, want true")
500 + if !resultUpdated {
501 + t.Fatal("applyRelayDiscoveryResponse() confirm updated = false, want true")
502 }
503 + advertisedDescriptors = server.relaySet.AdvertisedDescriptors()
504 advertisedURLs = advertisedURLs[:0]
450 - for _, descriptor := range server.peerRegistry.advertisedPeers() {
451 - if apiURL := strings.TrimSpace(descriptor.APIHTTPSAddr); apiURL != "" {
452 - advertisedURLs = append(advertisedURLs, apiURL)
505 + for _, descriptor := range advertisedDescriptors {
506 + if strings.TrimSpace(descriptor.APIHTTPSAddr) == "" {
507 + continue
508 }
509 + advertisedURLs = append(advertisedURLs, descriptor.APIHTTPSAddr)
510 }
511 + sort.Strings(advertisedURLs)
512 advertisedURLs, err = utils.ExcludeLocalRelayURLs(advertisedURLs...)
513 if err != nil {
514 t.Fatalf("ExcludeLocalRelayURLs() advertised second error = %v", err)
@@ -494,14 +551,14 @@ func TestServerStartHidesDiscoveryRoutesWhenDisabled(t *testing.T) {
551 }
552 })
553
497 - resp, err := client.Get("https://" + utils.HostPortOrLoopback(server.APIAddr()) + types.PathDiscovery + "?root_host=localhost&name=demo")
554 + resp, err := client.Get("https://" + utils.HostPortOrLoopback(server.APIAddr()) + types.PathDiscovery)
555 if err != nil {
499 - t.Fatalf("GET discovery resolve error = %v", err)
556 + t.Fatalf("GET relay discovery error = %v", err)
557 }
558 defer resp.Body.Close()
559
560 if resp.StatusCode != http.StatusNotFound {
504 - t.Fatalf("GET discovery resolve status = %d, want %d", resp.StatusCode, http.StatusNotFound)
561 + t.Fatalf("GET relay discovery status = %d, want %d", resp.StatusCode, http.StatusNotFound)
562 }
563 if server.DiscoveryEnabled() {
564 t.Fatal("DiscoveryEnabled() = true, want false without configured discovery service")
portal/wireguard/overlay.go
+3 -3
@@ -154,17 +154,17 @@ func (o *Overlay) Client() *http.Client {
154 }
155 }
156
157 -func (o *Overlay) Sync(selfRelayID string, snapshot map[string]types.PeerState) error {
157 +func (o *Overlay) Sync(selfRelayID string, snapshot map[string]types.RelayState) error {
158 if o == nil || o.stack == nil {
159 return nil
160 }
161 return o.stack.ApplyPeers(peersForSnapshot(selfRelayID, snapshot))
162 }
163
164 -func peersForSnapshot(selfRelayID string, snapshot map[string]types.PeerState) []types.DesiredPeer {
164 +func peersForSnapshot(selfRelayID string, snapshot map[string]types.RelayState) []types.DesiredPeer {
165 peers := make([]types.DesiredPeer, 0, len(snapshot))
166 for _, state := range snapshot {
167 - if state.State != types.PeerStateVerified && state.State != types.PeerStateAdvertised {
167 + if state.Expired {
168 continue
169 }
170 desc := state.Descriptor
sdk/expose.go
+130 -448
@@ -38,22 +38,14 @@ type Exposure struct {
38 accepted chan net.Conn
39 datagrams chan types.DatagramFrame
40
41 - mu sync.RWMutex
42 - knownRelayURLs []string
43 - bannedRelayURLs []string
44 - discoveryPins map[string]discoveryIdentity
45 - discoveryRelayIDsByURL map[string]string
46 - listeners map[string]*Listener
41 + relaySet *discovery.RelaySet
42 + listenerMu sync.RWMutex
43 + relayListeners map[string]*Listener
44
45 closeOnce sync.Once
46 connSeq atomic.Uint64
47 }
48
52 -type discoveryIdentity struct {
53 - apiURL string
54 - signerPublicKey string
55 -}
56 -
49 type ExposeConfig struct {
50 RelayURLs []string
51 Name string
@@ -94,96 +86,51 @@ func Expose(ctx context.Context, cfg ExposeConfig) (*Exposure, error) {
86
87 exposureCtx, cancel := context.WithCancel(ctx)
88 exposure := &Exposure{
97 - cancel: cancel,
98 - done: exposureCtx.Done(),
99 - name: cfg.Name,
100 - TargetAddr: targetAddr,
101 - UDPAddr: udpAddr,
102 - reverseToken: cfg.ReverseToken,
103 - udpEnabled: cfg.UDPEnabled,
104 - banMITM: cfg.BanMITM,
105 - metadata: cfg.Metadata.Copy(),
106 - ownerAddress: identity.Address,
107 - rootCAPEM: append([]byte(nil), cfg.RootCAPEM...),
108 - discoveryEnabled: cfg.Discovery,
109 - accepted: make(chan net.Conn, max(len(relayURLs)*defaultReadyTarget*2, 1)),
110 - datagrams: make(chan types.DatagramFrame, max(len(relayURLs)*32, 1)),
111 - discoveryPins: make(map[string]discoveryIdentity),
112 - discoveryRelayIDsByURL: make(map[string]string),
113 - listeners: make(map[string]*Listener, len(relayURLs)),
89 + cancel: cancel,
90 + done: exposureCtx.Done(),
91 + name: cfg.Name,
92 + TargetAddr: targetAddr,
93 + UDPAddr: udpAddr,
94 + reverseToken: cfg.ReverseToken,
95 + udpEnabled: cfg.UDPEnabled,
96 + banMITM: cfg.BanMITM,
97 + metadata: cfg.Metadata.Copy(),
98 + ownerAddress: identity.Address,
99 + rootCAPEM: append([]byte(nil), cfg.RootCAPEM...),
100 + discoveryEnabled: cfg.Discovery,
101 + accepted: make(chan net.Conn, max(len(relayURLs)*defaultReadyTarget*2, 1)),
102 + datagrams: make(chan types.DatagramFrame, max(len(relayURLs)*32, 1)),
103 + relaySet: discovery.NewRelaySet(),
104 + relayListeners: make(map[string]*Listener, len(relayURLs)),
105 }
106
107 if len(relayURLs) > 0 {
117 - if _, err := exposure.setRelayURLs(relayURLs, true); err != nil {
108 + exposure.relaySet.ReplaceKnownRelayURLs(relayURLs)
109 + if err := exposure.reconcileRelayListeners(true); err != nil {
110 _ = exposure.Close()
111 return nil, err
112 }
113 }
114
123 - if len(relayURLs) > 0 {
124 - go exposure.monitorStartupCounts()
125 - }
115 if exposure.discoveryEnabled {
127 - go exposure.runDiscoveryLoop(exposureCtx)
116 + go exposure.runRelayDiscoveryLoop(exposureCtx)
117 }
118 go func() {
119 <-exposure.done
120 _ = exposure.Close()
121 }()
122
134 - if len(relayURLs) > 0 {
135 - exposure.mu.RLock()
136 - activeRelayURLs, _ := exposure.activeSnapshotLocked()
137 - exposure.mu.RUnlock()
138 - log.Info().
139 - Str("release_version", types.ReleaseVersion).
140 - Int("relay_count", len(activeRelayURLs)).
141 - Strs("relays", activeRelayURLs).
142 - Msg("exposure relay started")
143 - }
144 -
123 return exposure, nil
124 }
125
126 func (e *Exposure) ActiveRelayURLs() []string {
149 - e.mu.RLock()
150 - defer e.mu.RUnlock()
151 -
152 - if len(e.knownRelayURLs) == 0 || len(e.listeners) == 0 {
127 + if e == nil {
128 return nil
129 }
155 -
156 - activeRelayURLs := make([]string, 0, len(e.listeners))
157 - for _, relayURL := range e.knownRelayURLs {
158 - if _, ok := e.listeners[relayURL]; ok {
159 - activeRelayURLs = append(activeRelayURLs, relayURL)
160 - }
161 - }
162 - if len(activeRelayURLs) == 0 {
130 + if e.relaySet == nil {
131 return nil
132 }
165 - return activeRelayURLs
166 -}
167 -
168 -func (e *Exposure) activeSnapshotLocked() ([]string, []*Listener) {
169 - if len(e.knownRelayURLs) == 0 || len(e.listeners) == 0 {
170 - return nil, nil
171 - }
172 -
173 - activeRelayURLs := make([]string, 0, len(e.listeners))
174 - listeners := make([]*Listener, 0, len(e.listeners))
175 - for _, relayURL := range e.knownRelayURLs {
176 - listener, ok := e.listeners[relayURL]
177 - if !ok {
178 - continue
179 - }
180 - activeRelayURLs = append(activeRelayURLs, relayURL)
181 - listeners = append(listeners, listener)
182 - }
183 - if len(activeRelayURLs) == 0 {
184 - return nil, nil
185 - }
186 - return activeRelayURLs, listeners
133 + return e.relaySet.ActiveRelayURLs()
134 }
135
136 func (e *Exposure) Accept() (net.Conn, error) {
@@ -217,9 +164,16 @@ func (e *Exposure) Addr() net.Addr {
164
165 func (e *Exposure) RunHTTP(ctx context.Context, handler http.Handler, localAddr string) error {
166 var relayListener net.Listener
220 - e.mu.RLock()
221 - _, activeListeners := e.activeSnapshotLocked()
222 - e.mu.RUnlock()
167 + e.listenerMu.RLock()
168 + activeListeners := make([]*Listener, 0, len(e.relayListeners))
169 + for _, relayURL := range e.relaySet.ActiveRelayURLs() {
170 + listener, ok := e.relayListeners[relayURL]
171 + if !ok {
172 + continue
173 + }
174 + activeListeners = append(activeListeners, listener)
175 + }
176 + e.listenerMu.RUnlock()
177 if len(activeListeners) > 0 {
178 relayListener = e
179 }
@@ -334,14 +288,14 @@ func (e *Exposure) Close() error {
288 e.cancel()
289 }
290
337 - e.mu.RLock()
338 - relayURLs := make([]string, 0, len(e.listeners))
339 - listeners := make([]*Listener, 0, len(e.listeners))
340 - for relayURL, listener := range e.listeners {
291 + e.listenerMu.RLock()
292 + relayURLs := make([]string, 0, len(e.relayListeners))
293 + listeners := make([]*Listener, 0, len(e.relayListeners))
294 + for relayURL, listener := range e.relayListeners {
295 relayURLs = append(relayURLs, relayURL)
296 listeners = append(listeners, listener)
297 }
344 - e.mu.RUnlock()
298 + e.listenerMu.RUnlock()
299
300 for _, listener := range listeners {
301 if listener != nil {
@@ -363,185 +317,88 @@ func (e *Exposure) Close() error {
317 return closeErr
318 }
319
366 -func (e *Exposure) setRelayURLs(relayURLs []string, failOnError bool) ([]string, error) {
367 - if len(relayURLs) == 0 {
368 - return nil, nil
369 - }
370 -
371 - e.mu.RLock()
372 - bannedRelayURLs := append([]string(nil), e.bannedRelayURLs...)
373 - e.mu.RUnlock()
374 - relayURLs = utils.FilterRelayURLs(append([]string(nil), relayURLs...), bannedRelayURLs)
375 -
376 - e.mu.Lock()
377 - existing := make(map[string]struct{}, len(e.knownRelayURLs))
378 - for _, relayURL := range e.knownRelayURLs {
379 - existing[relayURL] = struct{}{}
320 +func (e *Exposure) reconcileRelayListeners(failOnError bool) error {
321 + if e.relaySet == nil {
322 + e.relaySet = discovery.NewRelaySet()
323 }
381 - desired := make(map[string]struct{}, len(relayURLs))
382 - for _, relayURL := range relayURLs {
383 - desired[relayURL] = struct{}{}
324 + e.listenerMu.Lock()
325 + if e.relayListeners == nil {
326 + e.relayListeners = make(map[string]*Listener)
327 }
385 -
386 - added := make([]string, 0, len(relayURLs))
387 - missing := make([]string, 0)
388 - for _, relayURL := range relayURLs {
389 - if _, ok := existing[relayURL]; !ok {
390 - added = append(added, relayURL)
391 - }
392 - if _, ok := e.listeners[relayURL]; !ok {
393 - missing = append(missing, relayURL)
394 - }
328 + activeRelayURLs := e.relaySet.ActiveRelayURLs()
329 + currentRelayURLs := make([]string, 0, len(e.relayListeners))
330 + for relayURL := range e.relayListeners {
331 + currentRelayURLs = append(currentRelayURLs, relayURL)
332 }
396 - removed := make([]string, 0)
397 - staleListeners := make([]*Listener, 0)
398 - for relayURL, listener := range e.listeners {
399 - if _, ok := desired[relayURL]; ok {
400 - continue
401 - }
402 - removed = append(removed, relayURL)
403 - staleListeners = append(staleListeners, listener)
404 - delete(e.listeners, relayURL)
333 + missingRelayURLs := utils.FilterRelayURLs(activeRelayURLs, currentRelayURLs)
334 + staleRelayURLs := utils.FilterRelayURLs(currentRelayURLs, activeRelayURLs)
335 + staleListeners := make([]*Listener, 0, len(staleRelayURLs))
336 + for _, relayURL := range staleRelayURLs {
337 + staleListeners = append(staleListeners, e.relayListeners[relayURL])
338 + delete(e.relayListeners, relayURL)
339 }
406 - e.knownRelayURLs = append([]string(nil), relayURLs...)
407 - e.mu.Unlock()
340 + e.listenerMu.Unlock()
341
342 for i, listener := range staleListeners {
343 if listener == nil {
344 continue
345 }
346 if err := listener.Close(); err != nil && !errors.Is(err, net.ErrClosed) {
414 - log.Warn().Err(err).Str("relay_url", removed[i]).Msg("close stale relay listener")
415 - }
416 - }
417 - if len(removed) > 0 {
418 - log.Info().Int("removed_count", len(removed)).Strs("removed_relays", removed).Msg("stale relays removed from exposure")
419 - }
420 -
421 - for _, relayURL := range missing {
422 - listener, err := e.newListener(relayURL)
347 + log.Warn().Err(err).Str("relay_url", staleRelayURLs[i]).Msg("close stale relay listener")
348 + }
349 + }
350 + for _, relayURL := range missingRelayURLs {
351 + listener, err := NewListener(context.Background(), relayURL, ListenerConfig{
352 + Name: e.name,
353 + ReverseToken: e.reverseToken,
354 + UDPEnabled: e.udpEnabled,
355 + BanMITM: e.banMITM,
356 + Metadata: e.metadata.Copy(),
357 + RootCAPEM: append([]byte(nil), e.rootCAPEM...),
358 + ownerAddress: e.ownerAddress,
359 + relaySet: e.relaySet,
360 + })
361 if err != nil {
362 if failOnError {
425 - return nil, fmt.Errorf("listen %q: %w", relayURL, err)
363 + return fmt.Errorf("listen %q: %w", relayURL, err)
364 }
365 log.Warn().Err(err).Str("relay_url", relayURL).Msg("add relay listener")
366 continue
367 }
430 - e.installListener(relayURL, listener)
431 - }
432 - return added, nil
433 -}
368
435 -func (e *Exposure) pinDiscoverySelfDescriptor(targetURL string, desc types.RelayDescriptor) error {
436 - normalizedTargetURL, err := utils.NormalizeRelayURL(targetURL)
437 - if err != nil {
438 - return err
439 - }
440 - if desc.APIHTTPSAddr != normalizedTargetURL {
441 - return errors.New("descriptor api_https_addr does not match target url")
442 - }
443 - return e.pinDiscoveredDescriptor(desc)
444 -}
445 -
446 -func (e *Exposure) pinDiscoveredDescriptor(desc types.RelayDescriptor) error {
447 - relayID := strings.TrimSpace(desc.RelayID)
448 - apiURL := strings.TrimSpace(desc.APIHTTPSAddr)
449 - signerPublicKey := strings.TrimSpace(desc.SignerPublicKey)
450 - if relayID == "" {
451 - return errors.New("descriptor relay_id is required")
452 - }
453 - if apiURL == "" {
454 - return errors.New("descriptor api_https_addr is required")
455 - }
456 - if signerPublicKey == "" {
457 - return errors.New("descriptor signer_public_key is required")
458 - }
459 -
460 - e.mu.Lock()
461 - defer e.mu.Unlock()
462 - if e.discoveryPins == nil {
463 - e.discoveryPins = make(map[string]discoveryIdentity)
464 - }
465 - if e.discoveryRelayIDsByURL == nil {
466 - e.discoveryRelayIDsByURL = make(map[string]string)
467 - }
369 + select {
370 + case <-e.done:
371 + _ = listener.Close()
372 + continue
373 + default:
374 + }
375
469 - if pinned, ok := e.discoveryPins[relayID]; ok {
470 - if pinned.apiURL != apiURL {
471 - return errors.New("descriptor api_https_addr does not match pinned relay url")
376 + e.listenerMu.Lock()
377 + if e.relayListeners == nil {
378 + e.relayListeners = make(map[string]*Listener, 1)
379 }
473 - if pinned.signerPublicKey != signerPublicKey {
474 - return errors.New("descriptor signer_public_key does not match pinned signer")
380 + e.relaySet.MarkRelayUnreachable(relayURL)
381 + if _, exists := e.relayListeners[relayURL]; exists {
382 + e.listenerMu.Unlock()
383 + _ = listener.Close()
384 + continue
385 }
476 - }
477 - if pinnedRelayID, ok := e.discoveryRelayIDsByURL[apiURL]; ok && pinnedRelayID != relayID {
478 - return errors.New("descriptor relay_id does not match pinned relay url identity")
479 - }
386 + e.relayListeners[relayURL] = listener
387 + e.listenerMu.Unlock()
388
481 - e.discoveryPins[relayID] = discoveryIdentity{
482 - apiURL: apiURL,
483 - signerPublicKey: signerPublicKey,
389 + go e.runListenerAcceptLoop(listener)
390 }
485 - e.discoveryRelayIDsByURL[apiURL] = relayID
391 return nil
392 }
393
489 -func (e *Exposure) banRelayURL(relayURL string) {
490 - e.mu.Lock()
491 - e.knownRelayURLs = utils.RemoveRelayURL(e.knownRelayURLs, relayURL)
492 - e.bannedRelayURLs = utils.AppendUniqueRelayURL(e.bannedRelayURLs, relayURL)
493 - delete(e.listeners, relayURL)
494 - bannedRelayURLs := append([]string(nil), e.bannedRelayURLs...)
495 - e.mu.Unlock()
496 -
497 - log.Warn().
498 - Str("relay_url", relayURL).
499 - Strs("banned_relays", bannedRelayURLs).
500 - Msg("relay banned by mitm detection")
501 -}
502 -
503 -func (e *Exposure) newListener(relayURL string) (*Listener, error) {
504 - cfg := ListenerConfig{
505 - Name: e.name,
506 - ReverseToken: e.reverseToken,
507 - UDPEnabled: e.udpEnabled,
508 - BanMITM: e.banMITM,
509 - Metadata: e.metadata.Copy(),
510 - RootCAPEM: append([]byte(nil), e.rootCAPEM...),
511 - ownerAddress: e.ownerAddress,
512 - }
513 - return NewListener(context.Background(), relayURL, cfg)
514 -}
515 -
516 -func (e *Exposure) installListener(relayURL string, listener *Listener) {
394 +func (e *Exposure) runListenerAcceptLoop(listener *Listener) {
395 if listener == nil {
396 return
397 }
398
521 - shouldClose := false
522 - e.mu.Lock()
523 - select {
524 - case <-e.done:
525 - shouldClose = true
526 - default:
527 - if _, exists := e.listeners[relayURL]; exists {
528 - shouldClose = true
529 - } else {
530 - e.listeners[relayURL] = listener
531 - }
532 - }
533 - e.mu.Unlock()
534 -
535 - if shouldClose {
536 - _ = listener.Close()
537 - return
538 - }
539 -
540 - log.Info().Str("relay_url", relayURL).Msg("relay added to exposure")
541 - go e.runListenerAcceptLoop(listener)
399 + relayURL := listener.api.baseURL.String()
400 if e.udpEnabled {
401 go func() {
544 - relayURL := listener.api.baseURL.String()
402 for {
403 frame, err := listener.AcceptDatagram()
404 if err != nil {
@@ -569,25 +426,12 @@ func (e *Exposure) installListener(relayURL string, listener *Listener) {
426 }
427 }()
428 }
572 -}
573 -
574 -func (e *Exposure) runListenerAcceptLoop(listener *Listener) {
575 - if listener == nil {
576 - return
577 - }
578 -
579 - relayURL := listener.api.baseURL.String()
429 defer func() {
581 - if listener.StartupStatus() == listenerStatusBanned {
582 - e.banRelayURL(relayURL)
583 - return
584 - }
585 -
586 - e.mu.Lock()
587 - if current, ok := e.listeners[relayURL]; ok && current == listener {
588 - delete(e.listeners, relayURL)
430 + e.listenerMu.Lock()
431 + if current, ok := e.relayListeners[relayURL]; ok && current == listener {
432 + delete(e.relayListeners, relayURL)
433 }
590 - e.mu.Unlock()
434 + e.listenerMu.Unlock()
435 }()
436
437 for {
@@ -596,6 +440,7 @@ func (e *Exposure) runListenerAcceptLoop(listener *Listener) {
440 if listener.closed() || errors.Is(err, net.ErrClosed) {
441 return
442 }
443 + e.relaySet.MarkRelayFailure(relayURL, time.Now().UTC())
444 log.Warn().Err(err).Str("relay_url", relayURL).Msg("exposure listener accept failed")
445 return
446 }
@@ -659,9 +504,9 @@ func (e *Exposure) SendDatagram(frame types.DatagramFrame) error {
504 return net.ErrClosed
505 }
506
662 - e.mu.RLock()
663 - listener := e.listeners[frame.RelayURL]
664 - e.mu.RUnlock()
507 + e.listenerMu.RLock()
508 + listener := e.relayListeners[frame.RelayURL]
509 + e.listenerMu.RUnlock()
510 if listener == nil {
511 return net.ErrClosed
512 }
@@ -677,9 +522,12 @@ func (e *Exposure) WaitDatagramReady(ctx context.Context) ([]string, error) {
522 defer ticker.Stop()
523
524 for {
680 - e.mu.RLock()
681 - _, listeners := e.activeSnapshotLocked()
682 - e.mu.RUnlock()
525 + e.listenerMu.RLock()
526 + listeners := make([]*Listener, 0, len(e.relayListeners))
527 + for _, listener := range e.relayListeners {
528 + listeners = append(listeners, listener)
529 + }
530 + e.listenerMu.RUnlock()
531
532 addrs := make([]string, 0, len(listeners))
533 seen := make(map[string]struct{})
@@ -717,215 +565,49 @@ func (e *Exposure) WaitDatagramReady(ctx context.Context) ([]string, error) {
565 }
566 }
567
720 -func (e *Exposure) monitorStartupCounts() {
721 - ticker := time.NewTicker(time.Second)
722 - defer ticker.Stop()
723 - var (
724 - lastStatuses map[string]listenerStatus
725 - lastBannedCount = -1
726 - )
568 +const defaultDiscoveryInterval = 30 * time.Second
569
570 +func (e *Exposure) runRelayDiscoveryLoop(ctx context.Context) {
571 for {
729 - e.mu.RLock()
730 - _, listeners := e.activeSnapshotLocked()
731 - bannedCount := len(e.bannedRelayURLs)
732 - e.mu.RUnlock()
733 -
734 - currentStatuses := make(map[string]listenerStatus, len(listeners))
735 - readyCount := 0
736 - activated := make([]string, 0)
737 - deactivated := make([]string, 0)
572 + relayURLs := append([]string(nil), e.relaySet.ActiveRelayURLs()...)
573 + if len(relayURLs) > 0 {
574 + var discoveredRelayURLs []string
575
739 - for _, listener := range listeners {
740 - if listener == nil {
741 - continue
742 - }
743 - relayURL := listener.api.baseURL.String()
744 -
745 - status := listener.StartupStatus()
746 - currentStatuses[relayURL] = status
747 - if status == listenerStatusReady {
748 - readyCount++
749 - }
750 -
751 - if lastStatuses != nil {
752 - if prev, ok := lastStatuses[relayURL]; ok && prev != status {
753 - if status == listenerStatusReady {
754 - activated = append(activated, relayURL)
755 - } else {
756 - deactivated = append(deactivated, relayURL)
576 + for _, relayURL := range relayURLs {
577 + resp, err := discovery.DiscoverRelayDiscovery(ctx, relayURL, e.rootCAPEM, nil)
578 + if err != nil {
579 + if ctx.Err() != nil {
580 + return
581 }
582 + continue
583 }
759 - }
760 - }
584
762 - changed := lastStatuses == nil || bannedCount != lastBannedCount || len(currentStatuses) != len(lastStatuses)
763 - if !changed {
764 - for relayURL, status := range currentStatuses {
765 - if lastStatuses[relayURL] != status {
766 - changed = true
767 - break
768 - }
769 - }
770 - }
771 - if changed {
772 - inactiveCount := len(currentStatuses) - readyCount
773 - for relayURL, status := range lastStatuses {
774 - if _, ok := currentStatuses[relayURL]; ok || status != listenerStatusReady {
585 + now := time.Now().UTC()
586 + var descriptorRelayURLs []string
587 + descriptorRelayURLs, _, _, _, err = e.relaySet.ApplyRelayDiscoveryResponse(relayURL, relayURL, resp, now)
588 + if err != nil {
589 continue
590 }
777 - deactivated = append(deactivated, relayURL)
778 - }
779 -
780 - event := log.Info().
781 - Int("banned", bannedCount).
782 - Int("inactive", inactiveCount).
783 - Int("ready", readyCount)
784 - if len(activated) > 0 {
785 - event = event.Strs("activated", activated)
786 - }
787 - if len(deactivated) > 0 {
788 - event = event.Strs("deactivated", deactivated)
789 - }
790 - event.Msg("relay status")
791 - lastStatuses = currentStatuses
792 - lastBannedCount = bannedCount
793 - }
794 -
795 - select {
796 - case <-e.done:
797 - return
798 - case <-ticker.C:
799 - }
800 - }
801 -}
802 -
803 -const defaultDiscoveryInterval = 30 * time.Second
804 -
805 -func (e *Exposure) runDiscoveryLoop(ctx context.Context) {
806 - ticker := time.NewTicker(defaultDiscoveryInterval)
807 - defer ticker.Stop()
808 - discoveryFailed := false
809 -
810 - for {
811 - e.mu.RLock()
812 - relayURLs := append([]string(nil), e.knownRelayURLs...)
813 - e.mu.RUnlock()
814 - if len(relayURLs) == 0 {
815 - select {
816 - case <-ctx.Done():
817 - return
818 - case <-ticker.C:
819 - continue
820 - }
821 - }
822 -
823 - discoveredRelayURLs := append([]string(nil), relayURLs...)
824 - successCount := 0
825 - var discoveryErr error
826 - var warnErr error
591
828 - for _, relayURL := range relayURLs {
829 - resp, err := discovery.Discover(ctx, relayURL, types.DiscoverRequest{}, e.rootCAPEM, nil)
830 - if err != nil {
831 - if ctx.Err() != nil {
832 - return
592 + if len(discoveredRelayURLs) == 0 {
593 + discoveredRelayURLs = append([]string(nil), relayURLs...)
594 }
834 - discoveryErr = errors.Join(discoveryErr, fmt.Errorf("discover %q: %w", relayURL, err))
835 - continue
836 - }
837 -
838 - selfDescriptor, peerDescriptors, validateErr := discovery.ValidateResponse(resp, time.Now().UTC())
839 - if selfDescriptor.RelayID == "" || strings.TrimSpace(selfDescriptor.APIHTTPSAddr) == "" {
840 - discoveryErr = errors.Join(discoveryErr, fmt.Errorf("discover %q: missing self descriptor", relayURL))
841 - continue
842 - }
843 - if err := e.pinDiscoverySelfDescriptor(relayURL, selfDescriptor); err != nil {
844 - discoveryErr = errors.Join(discoveryErr, fmt.Errorf("discover %q: %w", relayURL, err))
845 - continue
846 - }
847 - if validateErr != nil {
848 - warnErr = errors.Join(warnErr, fmt.Errorf("discover %q: %w", relayURL, validateErr))
849 - }
850 -
851 - descriptorRelayURLs := make([]string, 0, 1+len(peerDescriptors))
852 - descriptorRelayURLs = append(descriptorRelayURLs, selfDescriptor.APIHTTPSAddr)
853 - for _, descriptor := range peerDescriptors {
854 - if err := e.pinDiscoveredDescriptor(descriptor); err != nil {
855 - warnErr = errors.Join(warnErr, fmt.Errorf("discover %q peer %q: %w", relayURL, descriptor.RelayID, err))
595 + discoveredRelayURLs, err = utils.MergeRelayURLs(discoveredRelayURLs, nil, descriptorRelayURLs)
596 + if err != nil {
597 continue
598 }
858 - if apiURL := strings.TrimSpace(descriptor.APIHTTPSAddr); apiURL != "" {
859 - descriptorRelayURLs = append(descriptorRelayURLs, apiURL)
860 - }
861 - }
862 -
863 - discoveredRelayURLs, err = utils.MergeRelayURLs(discoveredRelayURLs, nil, descriptorRelayURLs)
864 - if err != nil {
865 - discoveryErr = errors.Join(discoveryErr, fmt.Errorf("merge %q discovery relays: %w", relayURL, err))
866 - continue
599 }
868 - successCount++
869 - }
600
871 - switch {
872 - case successCount > 0:
873 - recovered := discoveryFailed
874 - discoveryFailed = false
875 - added, err := e.setRelayURLs(discoveredRelayURLs, false)
876 - if err != nil {
877 - log.Warn().
878 - Err(err).
879 - Int("relay_count", len(relayURLs)).
880 - Msg("relay discovery update failed")
881 - } else if recovered || len(added) > 0 || warnErr != nil || discoveryErr != nil {
882 - e.mu.RLock()
883 - totalKnownRelayCount := len(e.knownRelayURLs)
884 - e.mu.RUnlock()
885 - logErr := errors.Join(warnErr, discoveryErr)
886 - event := log.Info().
887 - Int("relay_count", len(relayURLs)).
888 - Int("discovered_count", successCount).
889 - Int("total_known_relay_count", totalKnownRelayCount)
890 - if recovered {
891 - event = event.Bool("recovered", true)
892 - }
893 - if len(added) > 0 {
894 - event = event.Int("added_count", len(added)).
895 - Strs("added_relays", added)
601 + if len(discoveredRelayURLs) > 0 {
602 + if e.relaySet == nil {
603 + e.relaySet = discovery.NewRelaySet()
604 }
897 - if logErr != nil {
898 - event = log.Warn().
899 - Err(logErr).
900 - Int("relay_count", len(relayURLs)).
901 - Int("discovered_count", successCount).
902 - Int("total_known_relay_count", totalKnownRelayCount)
903 - if recovered {
904 - event = event.Bool("recovered", true)
905 - }
906 - if len(added) > 0 {
907 - event = event.Int("added_count", len(added)).
908 - Strs("added_relays", added)
909 - }
910 - }
911 - event.Msg("relay discovery updated")
912 - }
913 - case ctx.Err() != nil:
914 - return
915 - default:
916 - if !discoveryFailed {
917 - log.Debug().
918 - Err(discoveryErr).
919 - Int("relay_count", len(relayURLs)).
920 - Msg("relay discovery failed")
605 + e.relaySet.ReplaceKnownRelayURLs(discoveredRelayURLs)
606 + _ = e.reconcileRelayListeners(false)
607 }
922 - discoveryFailed = true
608 }
924 -
925 - select {
926 - case <-ctx.Done():
609 + if !utils.SleepOrDone(ctx, defaultDiscoveryInterval) {
610 return
928 - case <-ticker.C:
611 }
612 }
613 }
sdk/expose_test.go
+88 -71
@@ -2,11 +2,41 @@ package sdk
2
3 import (
4 "net/url"
5 + "strings"
6 "testing"
7 + "time"
8
9 + "github.com/gosuda/portal/v2/portal/discovery"
10 "github.com/gosuda/portal/v2/types"
11 + "github.com/gosuda/portal/v2/utils"
12 )
13
14 +func mustSignedRelayDescriptor(t *testing.T, ownerPrivateKey, relayID, relayURL string) types.RelayDescriptor {
15 + t.Helper()
16 +
17 + identity, err := utils.ResolveSecp256k1Identity(ownerPrivateKey)
18 + if err != nil {
19 + t.Fatalf("ResolveSecp256k1Identity() error = %v", err)
20 + }
21 +
22 + now := time.Now().UTC()
23 + desc, err := discovery.SignedDescriptor(types.RelayDescriptor{
24 + RelayID: relayID,
25 + OwnerAddress: identity.Address,
26 + SignerPublicKey: identity.PublicKey,
27 + Sequence: uint64(now.UnixMilli()),
28 + Version: 1,
29 + IssuedAt: now,
30 + ExpiresAt: now.Add(time.Hour),
31 + APIHTTPSAddr: relayURL,
32 + StatusState: "healthy",
33 + }, identity.PrivateKey)
34 + if err != nil {
35 + t.Fatalf("SignedDescriptor() error = %v", err)
36 + }
37 + return desc
38 +}
39 +
40 func TestExposureBanRelayURLMovesRelay(t *testing.T) {
41 const (
42 relayA = "https://relay-a.example"
@@ -19,36 +49,35 @@ func TestExposureBanRelayURLMovesRelay(t *testing.T) {
49 }
50
51 listener := &Listener{
22 - api: &apiClient{baseURL: relayURL},
23 - startupStatus: listenerStatusBanned,
52 + api: &apiClient{baseURL: relayURL},
53 }
54
55 exposure := &Exposure{
27 - knownRelayURLs: []string{relayA, relayB},
28 - bannedRelayURLs: nil,
29 - listeners: map[string]*Listener{
30 - relayA: listener,
31 - relayB: {},
32 - },
56 + relaySet: discovery.NewRelaySet(),
57 + relayListeners: make(map[string]*Listener, 2),
58 + }
59 + exposure.relaySet.ReplaceKnownRelayURLs([]string{relayA, relayB})
60 + exposure.relayListeners = map[string]*Listener{
61 + relayA: listener,
62 + relayB: {},
63 }
64
35 - exposure.banRelayURL(relayA)
65 + exposure.relaySet.BanRelayURL(relayA, "mitm")
66 + exposure.listenerMu.Lock()
67 + delete(exposure.relayListeners, relayA)
68 + exposure.listenerMu.Unlock()
69
70 if got := exposure.ActiveRelayURLs(); len(got) != 1 || got[0] != relayB {
71 t.Fatalf("ActiveRelayURLs() = %v, want [%q]", got, relayB)
72 }
73
41 - exposure.mu.RLock()
42 - knownRelayURLs := append([]string(nil), exposure.knownRelayURLs...)
43 - bannedRelayURLs := append([]string(nil), exposure.bannedRelayURLs...)
44 - _, listenerExists := exposure.listeners[relayA]
45 - exposure.mu.RUnlock()
74 + knownRelayURLs := exposure.relaySet.ActiveRelayURLs()
75 + exposure.listenerMu.RLock()
76 + _, listenerExists := exposure.relayListeners[relayA]
77 + exposure.listenerMu.RUnlock()
78 if len(knownRelayURLs) != 1 || knownRelayURLs[0] != relayB {
79 t.Fatalf("knownRelayURLs = %v, want [%q]", knownRelayURLs, relayB)
80 }
49 - if len(bannedRelayURLs) != 1 || bannedRelayURLs[0] != relayA {
50 - t.Fatalf("bannedRelayURLs = %v, want [%q]", bannedRelayURLs, relayA)
51 - }
81 if listenerExists {
82 t.Fatal("banned relay listener still exists in exposure.listeners")
83 }
@@ -61,32 +90,25 @@ func TestExposureSetRelayURLsSkipsBannedRelay(t *testing.T) {
90 )
91
92 exposure := &Exposure{
64 - bannedRelayURLs: []string{relayB},
65 - listeners: map[string]*Listener{
66 - relayA: {},
67 - },
93 + relaySet: discovery.NewRelaySet(),
94 + relayListeners: make(map[string]*Listener, 1),
95 }
69 -
70 - added, err := exposure.setRelayURLs([]string{relayA, relayB}, false)
71 - if err != nil {
72 - t.Fatalf("setRelayURLs() error = %v", err)
96 + exposure.relaySet.BanRelayURL(relayB, "test")
97 + exposure.relayListeners = map[string]*Listener{
98 + relayA: {},
99 }
74 - if len(added) != 1 || added[0] != relayA {
75 - t.Fatalf("added relay urls = %v, want [%q]", added, relayA)
100 +
101 + exposure.relaySet.ReplaceKnownRelayURLs([]string{relayA, relayB})
102 + if err := exposure.reconcileRelayListeners(false); err != nil {
103 + t.Fatalf("reconcileRelayListeners() error = %v", err)
104 }
105 if got := exposure.ActiveRelayURLs(); len(got) != 1 || got[0] != relayA {
106 t.Fatalf("ActiveRelayURLs() = %v, want [%q]", got, relayA)
107 }
80 - exposure.mu.RLock()
81 - knownRelayURLs := append([]string(nil), exposure.knownRelayURLs...)
82 - bannedRelayURLs := append([]string(nil), exposure.bannedRelayURLs...)
83 - exposure.mu.RUnlock()
108 + knownRelayURLs := exposure.relaySet.ActiveRelayURLs()
109 if len(knownRelayURLs) != 1 || knownRelayURLs[0] != relayA {
110 t.Fatalf("knownRelayURLs = %v, want [%q]", knownRelayURLs, relayA)
111 }
87 - if len(bannedRelayURLs) != 1 || bannedRelayURLs[0] != relayB {
88 - t.Fatalf("bannedRelayURLs = %v, want [%q]", bannedRelayURLs, relayB)
89 - }
112 }
113
114 func TestExposureSetRelayURLsRemovesStaleListener(t *testing.T) {
@@ -106,25 +128,24 @@ func TestExposureSetRelayURLsRemovesStaleListener(t *testing.T) {
128
129 relayAClosed := make(chan struct{})
130 exposure := &Exposure{
109 - knownRelayURLs: []string{relayA, relayB},
110 - listeners: map[string]*Listener{
111 - relayA: {
112 - api: &apiClient{baseURL: relayAURL},
113 - cancel: func() { close(relayAClosed) },
114 - doneCh: relayAClosed,
115 - },
116 - relayB: {
117 - api: &apiClient{baseURL: relayBURL},
118 - },
131 + relaySet: discovery.NewRelaySet(),
132 + relayListeners: make(map[string]*Listener, 2),
133 + }
134 + exposure.relaySet.ReplaceKnownRelayURLs([]string{relayA, relayB})
135 + exposure.relayListeners = map[string]*Listener{
136 + relayA: {
137 + api: &apiClient{baseURL: relayAURL},
138 + cancel: func() { close(relayAClosed) },
139 + doneCh: relayAClosed,
140 + },
141 + relayB: {
142 + api: &apiClient{baseURL: relayBURL},
143 },
144 }
145
122 - added, err := exposure.setRelayURLs([]string{relayB}, false)
123 - if err != nil {
124 - t.Fatalf("setRelayURLs() error = %v", err)
125 - }
126 - if len(added) != 0 {
127 - t.Fatalf("added relay urls = %v, want empty", added)
146 + exposure.relaySet.ReplaceKnownRelayURLs([]string{relayB})
147 + if err := exposure.reconcileRelayListeners(false); err != nil {
148 + t.Fatalf("reconcileRelayListeners() error = %v", err)
149 }
150
151 select {
@@ -133,11 +154,11 @@ func TestExposureSetRelayURLsRemovesStaleListener(t *testing.T) {
154 t.Fatal("stale relay listener was not closed")
155 }
156
136 - exposure.mu.RLock()
137 - knownRelayURLs := append([]string(nil), exposure.knownRelayURLs...)
138 - _, relayAExists := exposure.listeners[relayA]
139 - _, relayBExists := exposure.listeners[relayB]
140 - exposure.mu.RUnlock()
157 + knownRelayURLs := exposure.relaySet.ActiveRelayURLs()
158 + exposure.listenerMu.RLock()
159 + _, relayAExists := exposure.relayListeners[relayA]
160 + _, relayBExists := exposure.relayListeners[relayB]
161 + exposure.listenerMu.RUnlock()
162 if len(knownRelayURLs) != 1 || knownRelayURLs[0] != relayB {
163 t.Fatalf("knownRelayURLs = %v, want [%q]", knownRelayURLs, relayB)
164 }
@@ -150,26 +171,22 @@ func TestExposureSetRelayURLsRemovesStaleListener(t *testing.T) {
171 }
172
173 func TestExposurePinDiscoveredDescriptorRejectsIdentityChange(t *testing.T) {
153 - exposure := &Exposure{}
154 - desc := types.RelayDescriptor{
155 - RelayID: "relay-a",
156 - APIHTTPSAddr: "https://relay-a.example",
157 - SignerPublicKey: "signer-a",
158 - }
174 + exposure := &Exposure{relaySet: discovery.NewRelaySet()}
175 + desc := mustSignedRelayDescriptor(t, strings.Repeat("11", 32), "relay-a", "https://relay-a.example")
176
160 - if err := exposure.pinDiscoverySelfDescriptor(desc.APIHTTPSAddr, desc); err != nil {
161 - t.Fatalf("pinDiscoverySelfDescriptor() error = %v", err)
177 + if _, _, _, _, err := exposure.relaySet.ApplyRelayDiscoveryResponse(desc.RelayID, desc.APIHTTPSAddr, types.DiscoveryResponse{Self: desc}, time.Now().UTC()); err != nil {
178 + t.Fatalf("ApplyRelayDiscoveryResponse() error = %v", err)
179 }
180
164 - changedSigner := desc
165 - changedSigner.SignerPublicKey = "signer-b"
166 - if err := exposure.pinDiscoveredDescriptor(changedSigner); err == nil {
167 - t.Fatal("pinDiscoveredDescriptor() error = nil, want pinned signer mismatch")
181 + changedSigner := mustSignedRelayDescriptor(t, strings.Repeat("12", 32), desc.RelayID, desc.APIHTTPSAddr)
182 + _, _, _, _, err := exposure.relaySet.ApplyRelayDiscoveryResponse(desc.RelayID, desc.APIHTTPSAddr, types.DiscoveryResponse{Self: changedSigner}, time.Now().UTC())
183 + if err == nil {
184 + t.Fatal("ApplyRelayDiscoveryResponse() error = nil, want pinned signer mismatch")
185 }
186
170 - changedURL := desc
171 - changedURL.APIHTTPSAddr = "https://relay-b.example"
172 - if err := exposure.pinDiscoveredDescriptor(changedURL); err == nil {
173 - t.Fatal("pinDiscoveredDescriptor() error = nil, want pinned relay url mismatch")
187 + changedURL := mustSignedRelayDescriptor(t, strings.Repeat("11", 32), desc.RelayID, "https://relay-b.example")
188 + _, _, _, _, err = exposure.relaySet.ApplyRelayDiscoveryResponse(desc.RelayID, "", types.DiscoveryResponse{Self: changedURL}, time.Now().UTC())
189 + if err == nil {
190 + t.Fatal("ApplyRelayDiscoveryResponse() error = nil, want pinned relay url mismatch")
191 }
192 }
sdk/listener.go
+36 -49
@@ -14,6 +14,7 @@ import (
14 "github.com/quic-go/quic-go"
15 "github.com/rs/zerolog/log"
16
17 + "github.com/gosuda/portal/v2/portal/discovery"
18 "github.com/gosuda/portal/v2/portal/keyless"
19 "github.com/gosuda/portal/v2/portal/transport"
20 "github.com/gosuda/portal/v2/types"
@@ -36,16 +37,9 @@ type ListenerConfig struct {
37 RetryCount int
38 RetryWait time.Duration
39 ownerAddress string
40 + relaySet *discovery.RelaySet
41 }
42
41 -type listenerStatus string
42 -
43 -const (
44 - listenerStatusInactive listenerStatus = "inactive"
45 - listenerStatusReady listenerStatus = "ready"
46 - listenerStatusBanned listenerStatus = "banned"
47 -)
48 -
43 type Listener struct {
44 api *apiClient
45 cancel context.CancelFunc
@@ -64,15 +58,15 @@ type Listener struct {
58 closeOnce sync.Once
59 registerOnce sync.Once
60
67 - banMITM bool
68 - mu sync.Mutex
69 - startupStatus listenerStatus
70 - leaseID string
71 - hostname string
72 - udpAddr string
73 - metadata types.LeaseMetadata
74 - tlsConfig *tls.Config
75 - tlsCloser io.Closer
61 + banMITM bool
62 + relaySet *discovery.RelaySet
63 + mu sync.Mutex
64 + leaseID string
65 + hostname string
66 + udpAddr string
67 + metadata types.LeaseMetadata
68 + tlsConfig *tls.Config
69 + tlsCloser io.Closer
70 }
71
72 // NewListener creates one relay listener and its dedicated relay transport for one relay URL.
@@ -92,17 +86,17 @@ func NewListener(ctx context.Context, relayURL string, cfg ListenerConfig) (*Lis
86 }
87
88 l := &Listener{
95 - doneCh: listenerCtx.Done(),
96 - cancel: cancel,
97 - api: api,
98 - registered: make(chan struct{}),
99 - startupStatus: listenerStatusInactive,
100 - retryCount: cfg.RetryCount,
101 - retryWait: retryWait,
102 - leaseTTL: leaseTTL,
103 - renewBefore: renewBefore,
104 - metadata: cfg.Metadata.Copy(),
105 - banMITM: cfg.BanMITM,
89 + doneCh: listenerCtx.Done(),
90 + cancel: cancel,
91 + api: api,
92 + registered: make(chan struct{}),
93 + retryCount: cfg.RetryCount,
94 + retryWait: retryWait,
95 + leaseTTL: leaseTTL,
96 + renewBefore: renewBefore,
97 + metadata: cfg.Metadata.Copy(),
98 + banMITM: cfg.BanMITM,
99 + relaySet: cfg.relaySet,
100 }
101 l.mitmManager = newMITMManager(listenerCtx, l)
102 l.stream = transport.NewClientStream(readyTarget, handshakeTimeout)
@@ -144,8 +138,8 @@ func (l *Listener) runStartup(ctx context.Context, readyTarget int) {
138 defer l.mu.Unlock()
139 return l.tlsConfig
140 },
147 - func() { l.setStartupStatus(listenerStatusReady) },
148 - func() { l.setStartupStatus(listenerStatusInactive) },
141 + func() { l.markReachable() },
142 + func() { l.markUnreachable() },
143 l.retryOrClose,
144 )
145 }
@@ -511,7 +505,7 @@ func (l *Listener) retryOrClose(ctx context.Context, operation string, err error
505 Logger()
506
507 if operation == "lease registration" {
514 - l.setStartupStatus(listenerStatusInactive)
508 + l.markUnreachable()
509 }
510
511 if l.retryCount > 0 && retries > l.retryCount {
@@ -554,35 +548,28 @@ func (l *Listener) closed() bool {
548 }
549 }
550
557 -func (l *Listener) setStartupStatus(status listenerStatus) {
551 +func (l *Listener) ban() {
552 if l == nil {
553 return
554 }
561 - l.mu.Lock()
562 - if l.startupStatus == listenerStatusBanned && status != listenerStatusBanned {
563 - l.mu.Unlock()
564 - return
555 + if l.relaySet != nil && l.api != nil && l.api.baseURL != nil {
556 + l.relaySet.BanRelayURL(l.api.baseURL.String(), "mitm")
557 }
566 - l.startupStatus = status
567 - l.mu.Unlock()
558 + _ = l.Close()
559 }
560
570 -func (l *Listener) ban() {
571 - if l == nil {
561 +func (l *Listener) markReachable() {
562 + if l == nil || l.relaySet == nil || l.api == nil || l.api.baseURL == nil {
563 return
564 }
574 - l.setStartupStatus(listenerStatusBanned)
575 - _ = l.Close()
565 + l.relaySet.MarkRelayReachable(l.api.baseURL.String(), time.Now().UTC())
566 }
567
578 -func (l *Listener) StartupStatus() listenerStatus {
579 - if l == nil {
580 - return listenerStatusInactive
568 +func (l *Listener) markUnreachable() {
569 + if l == nil || l.relaySet == nil || l.api == nil || l.api.baseURL == nil {
570 + return
571 }
582 -
583 - l.mu.Lock()
584 - defer l.mu.Unlock()
585 - return l.startupStatus
572 + l.relaySet.MarkRelayUnreachable(l.api.baseURL.String())
573 }
574
575 func (l *Listener) BanMITM() bool {
sdk/mitm_test.go
+15 -7
@@ -18,6 +18,7 @@ import (
18 "testing"
19 "time"
20
21 + "github.com/gosuda/portal/v2/portal/discovery"
22 "github.com/gosuda/portal/v2/types"
23 )
24
@@ -207,7 +208,8 @@ func TestMITMProbeDetectionBansListener(t *testing.T) {
208 }
209
210 listener := &Listener{
210 - api: &apiClient{baseURL: relayURL},
211 + api: &apiClient{baseURL: relayURL},
212 + relaySet: discovery.NewRelaySet(),
213 cancel: func() {
214 select {
215 case <-doneCh:
@@ -219,8 +221,9 @@ func TestMITMProbeDetectionBansListener(t *testing.T) {
221 registered: make(chan struct{}),
222 banMITM: true,
223 }
224 + listener.relaySet.ReplaceKnownRelayURLs([]string{relayURL.String()})
225 listener.mitmManager = newMITMManager(context.Background(), listener)
223 - listener.setStartupStatus(listenerStatusReady)
226 + listener.markReachable()
227
228 listener.mitmManager.logResult(MITMProbeReport{
229 RelayURL: relayURL.String(),
@@ -228,8 +231,10 @@ func TestMITMProbeDetectionBansListener(t *testing.T) {
231 Reason: types.MITMProbeReasonExporterMismatch,
232 }, nil)
233
231 - if status := listener.StartupStatus(); status != listenerStatusBanned {
232 - t.Fatalf("listener status = %q, want %q", status, listenerStatusBanned)
234 + for _, activeRelayURL := range listener.relaySet.ActiveRelayURLs() {
235 + if activeRelayURL == relayURL.String() {
236 + t.Fatal("relay still active after mitm detection")
237 + }
238 }
239 if !listener.closed() {
240 t.Fatal("listener.closed() = false, want true")
@@ -245,12 +250,14 @@ func TestMITMProbeDetectionWarnsWithoutBanningListener(t *testing.T) {
250
251 listener := &Listener{
252 api: &apiClient{baseURL: relayURL},
253 + relaySet: discovery.NewRelaySet(),
254 doneCh: doneCh,
255 registered: make(chan struct{}),
256 banMITM: false,
257 }
258 + listener.relaySet.ReplaceKnownRelayURLs([]string{relayURL.String()})
259 listener.mitmManager = newMITMManager(context.Background(), listener)
253 - listener.setStartupStatus(listenerStatusReady)
260 + listener.markReachable()
261
262 listener.mitmManager.logResult(MITMProbeReport{
263 RelayURL: relayURL.String(),
@@ -258,8 +265,9 @@ func TestMITMProbeDetectionWarnsWithoutBanningListener(t *testing.T) {
265 Reason: types.MITMProbeReasonExporterMismatch,
266 }, nil)
267
261 - if status := listener.StartupStatus(); status != listenerStatusReady {
262 - t.Fatalf("listener status = %q, want %q", status, listenerStatusReady)
268 + activeRelayURLs := listener.relaySet.ActiveRelayURLs()
269 + if len(activeRelayURLs) != 1 || activeRelayURLs[0] != relayURL.String() {
270 + t.Fatalf("ActiveRelayURLs() = %v, want [%q]", activeRelayURLs, relayURL.String())
271 }
272 if listener.closed() {
273 t.Fatal("listener.closed() = true, want false")
types/api.go
+5 -20
@@ -72,26 +72,11 @@ type RegisterResponse struct {
72 UDPEnabled bool `json:"udp_enabled,omitempty"`
73 }
74
75 -type DiscoverRequest struct {
76 - RootHost string `json:"root_host"`
77 - Name string `json:"name"`
78 -}
79 -
80 -type DiscoverResponse struct {
81 - ProtocolVersion uint32 `json:"protocol_version"`
82 - GeneratedAt time.Time `json:"generated_at"`
83 - Self RelayDescriptor `json:"self"`
84 - Peers []RelayDescriptor `json:"peers,omitempty"`
85 - Service *DiscoveredService `json:"service,omitempty"`
86 -}
87 -
88 -type DiscoveredService struct {
89 - Found bool `json:"found"`
90 - Name string `json:"name,omitempty"`
91 - Hostname string `json:"hostname,omitempty"`
92 - ExpiresAt time.Time `json:"expires_at,omitempty"`
93 - OwnerAddress string `json:"owner_address,omitempty"`
94 - RelayID string `json:"relay_id,omitempty"`
75 +type DiscoveryResponse struct {
76 + ProtocolVersion uint32 `json:"protocol_version"`
77 + GeneratedAt time.Time `json:"generated_at"`
78 + Self RelayDescriptor `json:"self"`
79 + Relays []RelayDescriptor `json:"relays,omitempty"`
80 }
81
82 type QUICControlMessage struct {
types/discovery.go
+8 -15
@@ -42,21 +42,14 @@ type RelayDescriptor struct {
42 DescriptorSignature string `json:"descriptor_signature"`
43 }
44
45 -type PeerLifecycleState string
46 -
47 -const (
48 - PeerStateKnown PeerLifecycleState = "known"
49 - PeerStateVerified PeerLifecycleState = "verified"
50 - PeerStateAdvertised PeerLifecycleState = "advertised"
51 - PeerStateExpired PeerLifecycleState = "expired"
52 -)
53 -
54 -type PeerState struct {
55 - Descriptor RelayDescriptor `json:"descriptor"`
56 - State PeerLifecycleState `json:"state"`
57 - FirstSeenAt time.Time `json:"first_seen_at"`
58 - LastSeenAt time.Time `json:"last_seen_at"`
59 - ConsecutiveFailures int `json:"consecutive_failures,omitempty"`
45 +type RelayState struct {
46 + Descriptor RelayDescriptor `json:"descriptor"`
47 + Bootstrap bool `json:"bootstrap,omitempty"`
48 + Advertised bool `json:"advertised,omitempty"`
49 + Expired bool `json:"expired,omitempty"`
50 + FirstSeenAt time.Time `json:"first_seen_at"`
51 + LastSeenAt time.Time `json:"last_seen_at"`
52 + ConsecutiveFailures int `json:"consecutive_failures,omitempty"`
53 }
54
55 type DesiredPeer struct {
utils/crypto.go
+10
@@ -1,6 +1,7 @@
1 package utils
2
3 import (
4 + "crypto/rand"
5 "crypto/sha256"
6 "encoding/base64"
7 "encoding/hex"
@@ -205,6 +206,15 @@ func NormalizeWireGuardPrivateKey(raw string) (string, error) {
206 return base64.StdEncoding.EncodeToString(key[:]), nil
207 }
208
209 +func GenerateWireGuardPrivateKey() (string, error) {
210 + var key [32]byte
211 + if _, err := rand.Read(key[:]); err != nil {
212 + return "", fmt.Errorf("generate wireguard private key: %w", err)
213 + }
214 + clampWireGuardPrivateKey(&key)
215 + return base64.StdEncoding.EncodeToString(key[:]), nil
216 +}
217 +
218 func WireGuardPublicKeyFromPrivate(raw string) (string, error) {
219 privateKey, err := decodeWireGuardKey(raw)
220 if err != nil {