refact: trust relay discovery via TLS instead of signed descriptors

Kim committed Mar 30, 2026 at 21:51 UTC 1ed3631bccea1a063a7022d853d996f6a1935cf7
13 files changed +217 -196
.env.example
+2
@@ -2,7 +2,9 @@
2 PORTAL_URL=https://localhost:4017
3 BOOTSTRAPS=https://localhost:4017
4 DISCOVERY=true
5 +OWNER_PRIVATE_KEY=
6 WIREGUARD_ENDPOINT=
7 +WIREGUARD_PRIVATE_KEY=
8
9 # Listener ports
10 API_PORT=4017
docker-compose.yml
+1
@@ -17,6 +17,7 @@ services:
17 PORTAL_URL: ${PORTAL_URL:-https://localhost:${API_PORT:-4017}}
18 BOOTSTRAPS: ${BOOTSTRAPS:-}
19 DISCOVERY: ${DISCOVERY:-true}
20 + OWNER_PRIVATE_KEY: ${OWNER_PRIVATE_KEY:-}
21 WIREGUARD_PRIVATE_KEY: ${WIREGUARD_PRIVATE_KEY:-}
22 DISCOVERY_PORT: ${DISCOVERY_PORT:-51820}
23 WIREGUARD_ENDPOINT: ${WIREGUARD_ENDPOINT:-}
docs/architecture.md
+1 -1
@@ -203,7 +203,7 @@ Wire format (`types/transport.go`): `[flowID uvarint][payload bytes]`
203 ## WireGuard Overlay and Discovery
204
205 - Discovery starts from bootstrap relay URLs over normal public HTTPS.
206 -- Each relay publishes a signed descriptor that may advertise:
206 +- Each relay publishes a descriptor over relay HTTPS that may advertise:
207 - `wireguard_public_key`
208 - `wireguard_endpoint`
209 - `overlay_ipv4`
docs/examples/nginx-proxy-multi-service/docker-compose.yaml
+1
@@ -70,6 +70,7 @@ services:
70 PORTAL_URL: ${PORTAL_URL:-https://portal.example.com}
71 API_PORT: ${API_PORT:-4017}
72 SNI_PORT: ${SNI_PORT:-4443}
73 + OWNER_PRIVATE_KEY: ${OWNER_PRIVATE_KEY:-}
74 WIREGUARD_PRIVATE_KEY: ${WIREGUARD_PRIVATE_KEY:-}
75 DISCOVERY_PORT: ${DISCOVERY_PORT:-51820}
76 WIREGUARD_ENDPOINT: ${WIREGUARD_ENDPOINT:-}
docs/examples/nginx-proxy/docker-compose.yaml
+1
@@ -64,6 +64,7 @@ services:
64 API_PORT: ${API_PORT:-4017}
65 # Use a non-443 port to avoid conflict with nginx on the host.
66 SNI_PORT: ${SNI_PORT:-4443}
67 + OWNER_PRIVATE_KEY: ${OWNER_PRIVATE_KEY:-}
68 WIREGUARD_PRIVATE_KEY: ${WIREGUARD_PRIVATE_KEY:-}
69 DISCOVERY_PORT: ${DISCOVERY_PORT:-51820}
70 WIREGUARD_ENDPOINT: ${WIREGUARD_ENDPOINT:-}
portal/api_server.go
+2 -4
@@ -128,10 +128,8 @@ func (s *Server) handleRelayDiscovery(w http.ResponseWriter, r *http.Request) {
128 strings.TrimSpace(s.wgConfig.Endpoint) != "" &&
129 strings.TrimSpace(s.wgConfig.OverlayIPv4) != ""
130
131 - self, err := discovery.SignedDescriptor(types.RelayDescriptor{
131 + self, err := discovery.NormalizeDescriptor(types.RelayDescriptor{
132 RelayID: s.cfg.PortalURL,
133 - OwnerAddress: s.ownerIdentity.Address,
134 - SignerPublicKey: s.ownerIdentity.PublicKey,
133 Sequence: uint64(now.UnixMilli()),
134 Version: 1,
135 IssuedAt: now,
@@ -145,7 +143,7 @@ func (s *Server) handleRelayDiscovery(w http.ResponseWriter, r *http.Request) {
143 WireGuardEndpoint: strings.TrimSpace(s.wgConfig.Endpoint),
144 OverlayIPv4: strings.TrimSpace(s.wgConfig.OverlayIPv4),
145 OverlayCIDRs: append([]string(nil), s.wgConfig.OverlayCIDRs...),
148 - }, s.ownerIdentity.PrivateKey)
146 + })
147 if err != nil {
148 utils.WriteAPIError(w, http.StatusInternalServerError, types.APIErrorCodeInternal, err.Error())
149 return
portal/discovery/discovery.go
-74
@@ -3,7 +3,6 @@ package discovery
3 import (
4 "context"
5 "crypto/tls"
6 - "encoding/json"
6 "errors"
7 "fmt"
8 "net/http"
@@ -20,12 +19,10 @@ const defaultRequestTimeout = 15 * time.Second
19
20 func NormalizeDescriptor(desc types.RelayDescriptor) (types.RelayDescriptor, error) {
21 desc.RelayID = strings.TrimSpace(desc.RelayID)
23 - desc.SignerPublicKey = strings.ToLower(strings.TrimSpace(desc.SignerPublicKey))
22 desc.APIHTTPSAddr = strings.TrimSpace(desc.APIHTTPSAddr)
23 desc.WireGuardPublicKey = strings.TrimSpace(desc.WireGuardPublicKey)
24 desc.WireGuardEndpoint = strings.TrimSpace(desc.WireGuardEndpoint)
25 desc.OverlayIPv4 = strings.TrimSpace(desc.OverlayIPv4)
28 - desc.DescriptorSignature = strings.TrimSpace(desc.DescriptorSignature)
26 if !desc.IssuedAt.IsZero() {
27 desc.IssuedAt = desc.IssuedAt.UTC()
28 }
@@ -43,13 +40,6 @@ func NormalizeDescriptor(desc types.RelayDescriptor) (types.RelayDescriptor, err
40 desc.RelayID = normalized
41 }
42 }
46 - if desc.OwnerAddress != "" {
47 - address, err := utils.NormalizeEVMAddress(desc.OwnerAddress)
48 - if err != nil {
49 - return types.RelayDescriptor{}, fmt.Errorf("normalize owner address: %w", err)
50 - }
51 - desc.OwnerAddress = address
52 - }
43 if len(desc.OverlayCIDRs) > 0 {
44 normalized, err := utils.NormalizeOverlayCIDRs(desc.OverlayCIDRs)
45 if err != nil {
@@ -66,32 +56,6 @@ func NormalizeDescriptor(desc types.RelayDescriptor) (types.RelayDescriptor, err
56 return desc, nil
57 }
58
69 -func SignDescriptor(desc types.RelayDescriptor, privateKeyHex string) (string, error) {
70 - normalized, err := NormalizeDescriptor(desc)
71 - if err != nil {
72 - return "", err
73 - }
74 - normalized.DescriptorSignature = ""
75 - payload, err := json.Marshal(normalized)
76 - if err != nil {
77 - return "", err
78 - }
79 - return utils.SignSHA256Secp256k1DER(payload, privateKeyHex)
80 -}
81 -
82 -func SignedDescriptor(desc types.RelayDescriptor, privateKeyHex string) (types.RelayDescriptor, error) {
83 - normalized, err := NormalizeDescriptor(desc)
84 - if err != nil {
85 - return types.RelayDescriptor{}, err
86 - }
87 - signature, err := SignDescriptor(normalized, privateKeyHex)
88 - if err != nil {
89 - return types.RelayDescriptor{}, err
90 - }
91 - normalized.DescriptorSignature = signature
92 - return normalized, nil
93 -}
94 -
59 func ValidateDescriptor(desc types.RelayDescriptor, now time.Time) (types.RelayDescriptor, error) {
60 normalized, err := NormalizeDescriptor(desc)
61 if err != nil {
@@ -120,14 +84,6 @@ func ValidateDescriptor(desc types.RelayDescriptor, now time.Time) (types.RelayD
84 case normalized.IssuedAt.After(normalized.ExpiresAt):
85 return types.RelayDescriptor{}, errors.New("issued_at must be before expires_at")
86 }
123 -
124 - derivedOwnerAddress, err := utils.AddressFromCompressedPublicKeyHex(normalized.SignerPublicKey)
125 - if err != nil {
126 - return types.RelayDescriptor{}, err
127 - }
128 - if normalized.OwnerAddress != derivedOwnerAddress {
129 - return types.RelayDescriptor{}, errors.New("owner_address does not match signer_public_key")
130 - }
87 if normalized.SupportsOverlayPeer {
88 if err := utils.ValidateWireGuardPublicKey(normalized.WireGuardPublicKey); err != nil {
89 return types.RelayDescriptor{}, err
@@ -139,17 +95,6 @@ func ValidateDescriptor(desc types.RelayDescriptor, now time.Time) (types.RelayD
95 return types.RelayDescriptor{}, err
96 }
97 }
142 -
143 - signature := normalized.DescriptorSignature
144 - normalized.DescriptorSignature = ""
145 - payload, err := json.Marshal(normalized)
146 - if err != nil {
147 - return types.RelayDescriptor{}, err
148 - }
149 - if err := utils.VerifySHA256Secp256k1DER(payload, normalized.SignerPublicKey, signature); err != nil {
150 - return types.RelayDescriptor{}, err
151 - }
152 - normalized.DescriptorSignature = signature
98 return normalized, nil
99 }
100
@@ -207,25 +152,6 @@ func ValidateDescriptorTarget(desc types.RelayDescriptor, targetRelayID, targetU
152 return nil
153 }
154
210 -// ValidateDescriptorMatch checks if a descriptor matches a pinned descriptor.
211 -func ValidateDescriptorMatch(desc, pinned types.RelayDescriptor) error {
212 - normalized, err := NormalizeDescriptor(desc)
213 - if err != nil {
214 - return err
215 - }
216 -
217 - if relayID := strings.TrimSpace(pinned.RelayID); relayID != "" && normalized.RelayID != relayID {
218 - return errors.New("descriptor relay_id does not match pinned relay id")
219 - }
220 - if apiURL := strings.TrimSpace(pinned.APIHTTPSAddr); apiURL != "" && normalized.APIHTTPSAddr != apiURL {
221 - return errors.New("descriptor api_https_addr does not match pinned relay url")
222 - }
223 - if signerKey := strings.ToLower(strings.TrimSpace(pinned.SignerPublicKey)); signerKey != "" && normalized.SignerPublicKey != signerKey {
224 - return errors.New("descriptor signer_public_key does not match pinned signer")
225 - }
226 - return nil
227 -}
228 -
155 func DiscoverRelayDiscovery(ctx context.Context, baseURL string, rootCAPEM []byte, httpClient *http.Client) (types.DiscoveryResponse, error) {
156 resp, err := doGET[types.DiscoveryResponse](ctx, baseURL, types.PathDiscovery, nil, rootCAPEM, httpClient)
157 if err != nil {
portal/discovery/relayset.go
+43 -61
@@ -44,8 +44,8 @@ type RelaySummary struct {
44 Unreachable int
45 }
46
47 -// RelaySet owns the shared relay discovery view: known relay URLs, pinned relay
48 -// descriptors, the latest validated descriptor seen for each relay, and common
47 +// RelaySet owns the shared relay discovery view: known relay URLs, stable relay
48 +// id/url mappings, the latest validated descriptor seen for each relay, and common
49 // process-local relay state such as ban/reachability/failure tracking.
50 //
51 // Runtime-specific policy such as bootstrap classification, relay lifecycle, or
@@ -53,7 +53,6 @@ type RelaySummary struct {
53 type RelaySet struct {
54 mu sync.RWMutex
55 knownRelayURLs []string
56 - pinnedByRelayID map[string]types.RelayDescriptor
56 relayIDsByURL map[string]string
57 relays map[string]RelayView
58 localByURL map[string]RelayLocalState
@@ -64,10 +63,9 @@ type RelaySet struct {
63
64 func NewRelaySet() *RelaySet {
65 return &RelaySet{
67 - pinnedByRelayID: make(map[string]types.RelayDescriptor),
68 - relayIDsByURL: make(map[string]string),
69 - relays: make(map[string]RelayView),
70 - localByURL: make(map[string]RelayLocalState),
66 + relayIDsByURL: make(map[string]string),
67 + relays: make(map[string]RelayView),
68 + localByURL: make(map[string]RelayLocalState),
69 }
70 }
71
@@ -126,7 +124,21 @@ func (s *RelaySet) ActiveRelayURLs() []string {
124 return out
125 }
126
127 +func relayExpiredAt(view RelayView, state RelayLocalState, now time.Time) bool {
128 + if state.Expired {
129 + return true
130 + }
131 + if view.Descriptor.ExpiresAt.IsZero() {
132 + return false
133 + }
134 + if now.IsZero() {
135 + now = time.Now().UTC()
136 + }
137 + return !view.Descriptor.ExpiresAt.After(now)
138 +}
139 +
140 func (s *RelaySet) logStatusChange() {
141 + now := time.Now().UTC()
142 var currentReachable map[string]bool
143 trackedRelayURLs := s.trackedRelayURLs()
144 if len(trackedRelayURLs) > 0 {
@@ -140,6 +152,9 @@ func (s *RelaySet) logStatusChange() {
152 for _, relayURL := range trackedRelayURLs {
153 summary.Known++
154 state := s.localByURL[relayURL]
155 + relayID := s.relayIDsByURL[relayURL]
156 + view, ok := s.relays[relayID]
157 + expired := ok && relayExpiredAt(view, state, now) || !ok && state.Expired
158 if state.Banned {
159 summary.Banned++
160 continue
@@ -147,10 +162,10 @@ func (s *RelaySet) logStatusChange() {
162 if state.Bootstrap {
163 summary.Bootstrap++
164 }
150 - if state.Advertised {
165 + if state.Advertised && !expired {
166 summary.Advertised++
167 }
153 - if state.Expired {
168 + if expired {
169 summary.Expired++
170 }
171 if state.Reachable {
@@ -158,9 +173,7 @@ func (s *RelaySet) logStatusChange() {
173 } else {
174 summary.Unreachable++
175 }
161 - relayID := s.relayIDsByURL[relayURL]
162 - view, ok := s.relays[relayID]
163 - if ok && !state.Bootstrap && !state.Expired && view.Descriptor.SupportsOverlayPeer {
176 + if ok && !state.Bootstrap && !expired && view.Descriptor.SupportsOverlayPeer {
177 summary.Syncable++
178 }
179 }
@@ -347,10 +360,11 @@ func (s *RelaySet) AdvertisedDescriptors() []types.RelayDescriptor {
360 return nil
361 }
362
363 + now := time.Now().UTC()
364 out := make([]types.RelayDescriptor, 0, len(s.relays))
365 for _, view := range s.relays {
366 state := s.localByURL[view.Descriptor.APIHTTPSAddr]
353 - if !state.Advertised || state.Expired || strings.TrimSpace(view.Descriptor.APIHTTPSAddr) == "" {
367 + if !state.Advertised || relayExpiredAt(view, state, now) || strings.TrimSpace(view.Descriptor.APIHTTPSAddr) == "" {
368 continue
369 }
370 out = append(out, view.Descriptor)
@@ -374,10 +388,11 @@ func (s *RelaySet) SyncableDescriptors() []types.RelayDescriptor {
388 return nil
389 }
390
391 + now := time.Now().UTC()
392 out := make([]types.RelayDescriptor, 0, len(s.relays))
393 for _, view := range s.relays {
394 state := s.localByURL[view.Descriptor.APIHTTPSAddr]
380 - if state.Bootstrap || state.Expired || !view.Descriptor.SupportsOverlayPeer {
395 + if state.Bootstrap || relayExpiredAt(view, state, now) || !view.Descriptor.SupportsOverlayPeer {
396 continue
397 }
398 out = append(out, view.Descriptor)
@@ -401,6 +416,7 @@ func (s *RelaySet) Snapshot() map[string]types.RelayState {
416 return nil
417 }
418
419 + now := time.Now().UTC()
420 snapshot := make(map[string]types.RelayState, len(s.relays))
421 for relayID, view := range s.relays {
422 localState := s.localByURL[view.Descriptor.APIHTTPSAddr]
@@ -408,7 +424,7 @@ func (s *RelaySet) Snapshot() map[string]types.RelayState {
424 Descriptor: view.Descriptor,
425 Bootstrap: localState.Bootstrap,
426 Advertised: localState.Advertised,
411 - Expired: localState.Expired,
427 + Expired: relayExpiredAt(view, localState, now),
428 FirstSeenAt: view.FirstSeenAt,
429 LastSeenAt: view.LastSeenAt,
430 ConsecutiveFailures: localState.ConsecutiveFailures,
@@ -444,20 +460,6 @@ func (s *RelaySet) ReplaceKnownRelayURLs(relayURLs []string) {
460 s.knownRelayURLs = append([]string(nil), filtered...)
461 }
462
447 -func (s *RelaySet) pinTarget(targetRelayID, targetURL string, desc types.RelayDescriptor) error {
448 - if s == nil {
449 - return nil
450 - }
451 - if err := ValidateDescriptorTarget(desc, targetRelayID, targetURL); err != nil {
452 - return err
453 - }
454 - if err := s.matchPinned(desc); err != nil {
455 - return err
456 - }
457 - s.pin(desc)
458 - return nil
459 -}
460 -
463 func (s *RelaySet) registerDescriptor(desc types.RelayDescriptor, now time.Time) (string, bool, bool, error) {
464 if s == nil {
465 return "", false, false, nil
@@ -466,10 +468,15 @@ func (s *RelaySet) registerDescriptor(desc types.RelayDescriptor, now time.Time)
468 if err != nil {
469 return "", false, false, err
470 }
469 - if err := s.matchPinned(normalized); err != nil {
470 - return "", false, false, err
471 + if current, ok := s.relays[normalized.RelayID]; ok {
472 + currentURL := strings.TrimSpace(current.Descriptor.APIHTTPSAddr)
473 + if currentURL != "" && currentURL != normalized.APIHTTPSAddr {
474 + return "", false, false, errors.New("descriptor api_https_addr does not match known relay url")
475 + }
476 + }
477 + if knownRelayID, ok := s.relayIDsByURL[normalized.APIHTTPSAddr]; ok && knownRelayID != normalized.RelayID {
478 + return "", false, false, errors.New("descriptor relay_id does not match known relay")
479 }
472 - s.pin(normalized)
480
481 if now.IsZero() {
482 now = time.Now().UTC()
@@ -482,11 +489,12 @@ func (s *RelaySet) registerDescriptor(desc types.RelayDescriptor, now time.Time)
489 view.FirstSeenAt = now
490 }
491 previousDescriptor := view.Descriptor
485 - view.Descriptor = desc
492 + view.Descriptor = normalized
493 view.LastSeenAt = now
494 s.relays[relayID] = view
495 + s.relayIDsByURL[normalized.APIHTTPSAddr] = relayID
496
489 - changed := added || !reflect.DeepEqual(previousDescriptor, desc)
497 + changed := added || !reflect.DeepEqual(previousDescriptor, normalized)
498 return relayID, added, changed, nil
499 }
500
@@ -516,7 +524,7 @@ func (s *RelaySet) applyDiscoveryDescriptors(targetRelayID, targetURL string, se
524 if now.IsZero() {
525 now = time.Now().UTC()
526 }
519 - if err := s.pinTarget(targetRelayID, targetURL, selfDescriptor); err != nil {
527 + if err := ValidateDescriptorTarget(selfDescriptor, targetRelayID, targetURL); err != nil {
528 return false, 0, err
529 }
530
@@ -715,29 +723,3 @@ func (s *RelaySet) RecordDiscoveryFailure(relayID, relayURL string, err error, r
723 }
724 return false, "", localState.ConsecutiveFailures
725 }
718 -
719 -func (s *RelaySet) matchPinned(desc types.RelayDescriptor) error {
720 - if s == nil {
721 - return nil
722 - }
723 - if pinned, ok := s.pinnedByRelayID[desc.RelayID]; ok {
724 - if err := ValidateDescriptorMatch(desc, pinned); err != nil {
725 - return err
726 - }
727 - }
728 - if pinnedRelayID, ok := s.relayIDsByURL[desc.APIHTTPSAddr]; ok && pinnedRelayID != desc.RelayID {
729 - return ValidateDescriptorMatch(desc, types.RelayDescriptor{
730 - RelayID: pinnedRelayID,
731 - APIHTTPSAddr: desc.APIHTTPSAddr,
732 - })
733 - }
734 - return nil
735 -}
736 -
737 -func (s *RelaySet) pin(desc types.RelayDescriptor) {
738 - if s == nil {
739 - return
740 - }
741 - s.pinnedByRelayID[desc.RelayID] = desc
742 - s.relayIDsByURL[desc.APIHTTPSAddr] = desc.RelayID
743 -}
portal/server_test.go
+62 -15
@@ -4,6 +4,7 @@ import (
4 "context"
5 "crypto/tls"
6 "encoding/json"
7 + "io"
8 "net"
9 "net/http"
10 "reflect"
@@ -18,14 +19,9 @@ import (
19 "github.com/gosuda/portal/v2/utils"
20 )
21
21 -func mustSignedRelayDescriptor(t *testing.T, ownerPrivateKey, relayURL string) types.RelayDescriptor {
22 +func mustRelayDescriptor(t *testing.T, relayURL string) types.RelayDescriptor {
23 t.Helper()
24
24 - identity, err := utils.ResolveSecp256k1Identity(ownerPrivateKey)
25 - if err != nil {
26 - t.Fatalf("ResolveSecp256k1Identity() error = %v", err)
27 - }
28 -
25 now := time.Now().UTC()
26 wireGuardPrivateKey, err := utils.NormalizeWireGuardPrivateKey(strings.Repeat("44", 32))
27 if err != nil {
@@ -39,10 +35,8 @@ func mustSignedRelayDescriptor(t *testing.T, ownerPrivateKey, relayURL string) t
35 if err != nil {
36 t.Fatalf("DeriveWireGuardOverlayIPv4() error = %v", err)
37 }
42 - desc, err := discovery.SignedDescriptor(types.RelayDescriptor{
38 + desc, err := discovery.NormalizeDescriptor(types.RelayDescriptor{
39 RelayID: relayURL,
44 - OwnerAddress: identity.Address,
45 - SignerPublicKey: identity.PublicKey,
40 Sequence: uint64(now.UnixMilli()),
41 Version: 1,
42 IssuedAt: now,
@@ -53,9 +47,9 @@ func mustSignedRelayDescriptor(t *testing.T, ownerPrivateKey, relayURL string) t
47 OverlayIPv4: overlayIPv4,
48 SupportsTCP: true,
49 SupportsOverlayPeer: true,
56 - }, identity.PrivateKey)
50 + })
51 if err != nil {
58 - t.Fatalf("SignedDescriptor() error = %v", err)
52 + t.Fatalf("NormalizeDescriptor() error = %v", err)
53 }
54 return desc
55 }
@@ -147,6 +141,61 @@ func TestServerStartInitializesLocalACMEAndSigner(t *testing.T) {
141 }
142 }
143
144 +func TestServerStartDiscoveryOmitsOwnerIdentityFields(t *testing.T) {
145 + t.Parallel()
146 +
147 + server, err := NewServer(ServerConfig{
148 + PortalURL: "https://localhost:4017",
149 + ACME: acme.Config{KeyDir: t.TempDir()},
150 + APIListenAddr: "127.0.0.1:0",
151 + SNIListenAddr: "127.0.0.1:0",
152 + DiscoveryEnabled: true,
153 + })
154 + if err != nil {
155 + t.Fatalf("NewServer() error = %v", err)
156 + }
157 +
158 + ctx, cancel := context.WithCancel(context.Background())
159 + defer cancel()
160 +
161 + if err := server.Start(ctx, nil); err != nil {
162 + t.Fatalf("Start() error = %v", err)
163 + }
164 +
165 + client := &http.Client{
166 + Transport: &http.Transport{
167 + TLSClientConfig: &tls.Config{InsecureSkipVerify: true},
168 + },
169 + }
170 + t.Cleanup(func() {
171 + client.CloseIdleConnections()
172 + cancel()
173 + if err := server.Wait(); err != nil {
174 + t.Fatalf("Wait() error = %v", err)
175 + }
176 + })
177 +
178 + resp, err := client.Get("https://" + utils.HostPortOrLoopback(server.APIAddr()) + types.PathDiscovery)
179 + if err != nil {
180 + t.Fatalf("GET /discovery error = %v", err)
181 + }
182 + defer resp.Body.Close()
183 +
184 + if resp.StatusCode != http.StatusOK {
185 + t.Fatalf("GET /discovery status = %d, want %d", resp.StatusCode, http.StatusOK)
186 + }
187 +
188 + body, err := io.ReadAll(resp.Body)
189 + if err != nil {
190 + t.Fatalf("read /discovery response: %v", err)
191 + }
192 + for _, key := range []string{"owner_address", "signer_public_key", "descriptor_signature"} {
193 + if strings.Contains(string(body), key) {
194 + t.Fatalf("/discovery body = %q, want %q omitted", string(body), key)
195 + }
196 + }
197 +}
198 +
199 func TestServerStartRejectsMismatchedACMEBaseDomain(t *testing.T) {
200 t.Parallel()
201
@@ -365,11 +414,9 @@ func TestServerUpsertDiscoverySeedURLsSkipsLocalRelayHosts(t *testing.T) {
414 func TestServerRecordVerifiedDiscoveryPeerRequiresDirectConfirmation(t *testing.T) {
415 t.Parallel()
416
368 - ownerPrivateKey := strings.Repeat("11", 32)
417 server, err := NewServer(ServerConfig{
418 PortalURL: "https://portal.example.com",
419 Bootstraps: []string{"https://bootstrap.example.com"},
372 - OwnerPrivateKey: ownerPrivateKey,
420 WireGuardPrivateKey: strings.Repeat("24", 32),
421 DiscoveryPort: 41023,
422 DiscoveryEnabled: true,
@@ -378,8 +425,8 @@ func TestServerRecordVerifiedDiscoveryPeerRequiresDirectConfirmation(t *testing.
425 t.Fatalf("NewServer() error = %v", err)
426 }
427
381 - bootstrapDesc := mustSignedRelayDescriptor(t, ownerPrivateKey, "https://bootstrap.example.com")
382 - relayADesc := mustSignedRelayDescriptor(t, ownerPrivateKey, "https://relay-a.example.com")
428 + bootstrapDesc := mustRelayDescriptor(t, "https://bootstrap.example.com")
429 + relayADesc := mustRelayDescriptor(t, "https://relay-a.example.com")
430
431 applyDiscovery := func(targetRelayID, targetURL string, resp types.DiscoveryResponse, requireSelfOverlay bool) (bool, int, error, error) {
432 now := time.Now().UTC()
portal/wireguard/stack.go
+22 -7
@@ -32,8 +32,9 @@ type stack struct {
32 net *netstack.Net
33 overlayIP netip.Addr
34
35 - mu sync.Mutex
36 - closed bool
35 + mu sync.Mutex
36 + closed bool
37 + peerEndpoints map[string]string
38 }
39
40 func newStack(cfg Config) (*stack, error) {
@@ -78,9 +79,10 @@ func newStack(cfg Config) (*stack, error) {
79 }
80
81 return &stack{
81 - device: wgDevice,
82 - net: network,
83 - overlayIP: overlayIP,
82 + device: wgDevice,
83 + net: network,
84 + overlayIP: overlayIP,
85 + peerEndpoints: map[string]string{},
86 }, nil
87 }
88
@@ -127,6 +129,7 @@ func (s *stack) ApplyPeers(peers []types.DesiredPeer) error {
129 var builder strings.Builder
130 builder.WriteString("replace_peers=true\n")
131 var warnErr error
132 + nextPeerEndpoints := map[string]string{}
133
134 for _, peer := range peers {
135 publicKeyHex, err := utils.WireGuardKeyHex(peer.WireGuardPublicKey)
@@ -138,8 +141,16 @@ func (s *stack) ApplyPeers(peers []types.DesiredPeer) error {
141 if endpoint := strings.TrimSpace(peer.WireGuardEndpoint); endpoint != "" {
142 resolvedEndpoint, err = resolvePeerEndpoint(endpoint)
143 if err != nil {
141 - warnErr = errors.Join(warnErr, fmt.Errorf("resolve peer %q endpoint: %w", peer.RelayID, err))
142 - continue
144 + s.mu.Lock()
145 + currentEndpoint := s.peerEndpoints[publicKeyHex]
146 + s.mu.Unlock()
147 + if currentEndpoint != "" {
148 + warnErr = errors.Join(warnErr, fmt.Errorf("resolve peer %q endpoint: %w; using current endpoint %q", peer.RelayID, err, currentEndpoint))
149 + resolvedEndpoint = currentEndpoint
150 + } else {
151 + warnErr = errors.Join(warnErr, fmt.Errorf("resolve peer %q endpoint: %w", peer.RelayID, err))
152 + continue
153 + }
154 }
155 }
156
@@ -150,6 +161,7 @@ func (s *stack) ApplyPeers(peers []types.DesiredPeer) error {
161 builder.WriteString("endpoint=")
162 builder.WriteString(resolvedEndpoint)
163 builder.WriteByte('\n')
164 + nextPeerEndpoints[publicKeyHex] = resolvedEndpoint
165 }
166
167 allowedIPs := utils.NormalizeIPPrefixes(peer.AllowedIPs)
@@ -168,6 +180,9 @@ func (s *stack) ApplyPeers(peers []types.DesiredPeer) error {
180 if err := s.device.IpcSet(builder.String()); err != nil {
181 return err
182 }
183 + s.mu.Lock()
184 + s.peerEndpoints = nextPeerEndpoints
185 + s.mu.Unlock()
186 return warnErr
187 }
188
portal/wireguard/stack_test.go
+68
@@ -6,6 +6,7 @@ import (
6 "strings"
7 "testing"
8
9 + "github.com/gosuda/portal/v2/types"
10 "github.com/gosuda/portal/v2/utils"
11 )
12
@@ -97,6 +98,73 @@ func TestResolvePeerEndpointResolvesHostname(t *testing.T) {
98 }
99 }
100
101 +func TestApplyPeersKeepsCurrentEndpointOnResolveFailure(t *testing.T) {
102 + t.Parallel()
103 +
104 + privateKey, err := utils.NormalizeWireGuardPrivateKey("3333333333333333333333333333333333333333333333333333333333333333")
105 + if err != nil {
106 + t.Fatalf("NormalizeWireGuardPrivateKey() error = %v", err)
107 + }
108 +
109 + port := reserveUDPPort(t)
110 + stack, err := newStack(Config{
111 + PrivateKey: privateKey,
112 + Endpoint: net.JoinHostPort("127.0.0.1", port),
113 + OverlayIPv4: "10.77.0.1",
114 + })
115 + if err != nil {
116 + t.Fatalf("newStack() error = %v", err)
117 + }
118 + t.Cleanup(func() {
119 + if err := stack.Close(); err != nil {
120 + t.Fatalf("Close() error = %v", err)
121 + }
122 + })
123 +
124 + peerPrivateKey, err := utils.NormalizeWireGuardPrivateKey("4444444444444444444444444444444444444444444444444444444444444444")
125 + if err != nil {
126 + t.Fatalf("NormalizeWireGuardPrivateKey() error = %v", err)
127 + }
128 + peerPublicKey, err := utils.WireGuardPublicKeyFromPrivate(peerPrivateKey)
129 + if err != nil {
130 + t.Fatalf("WireGuardPublicKeyFromPrivate() error = %v", err)
131 + }
132 + peerPublicKeyHex, err := utils.WireGuardKeyHex(peerPublicKey)
133 + if err != nil {
134 + t.Fatalf("WireGuardKeyHex() error = %v", err)
135 + }
136 +
137 + peer := types.DesiredPeer{
138 + RelayID: "https://peer.example.com",
139 + WireGuardPublicKey: peerPublicKey,
140 + WireGuardEndpoint: "127.0.0.1:51820",
141 + AllowedIPs: []string{"10.77.0.2/32"},
142 + }
143 + if err := stack.ApplyPeers([]types.DesiredPeer{peer}); err != nil {
144 + t.Fatalf("ApplyPeers() initial error = %v", err)
145 + }
146 +
147 + peer.WireGuardEndpoint = "peer.invalid:51820"
148 + err = stack.ApplyPeers([]types.DesiredPeer{peer})
149 + if err == nil {
150 + t.Fatal("ApplyPeers() warning error = nil, want resolve warning")
151 + }
152 + if !strings.Contains(err.Error(), "using current endpoint") {
153 + t.Fatalf("ApplyPeers() warning = %q, want current endpoint fallback", err)
154 + }
155 +
156 + config, err := stack.device.IpcGet()
157 + if err != nil {
158 + t.Fatalf("IpcGet() error = %v", err)
159 + }
160 + if !strings.Contains(config, "public_key="+peerPublicKeyHex+"\n") {
161 + t.Fatalf("IpcGet() = %q, want peer public key %q", config, peerPublicKeyHex)
162 + }
163 + if !strings.Contains(config, "endpoint=127.0.0.1:51820\n") {
164 + t.Fatalf("IpcGet() = %q, want endpoint %q", config, "127.0.0.1:51820")
165 + }
166 +}
167 +
168 func reserveUDPPort(t *testing.T) string {
169 t.Helper()
170
sdk/expose_test.go
+14 -29
@@ -2,36 +2,27 @@ package sdk
2
3 import (
4 "net/url"
5 - "strings"
5 "testing"
6 "time"
7
8 "github.com/gosuda/portal/v2/portal/discovery"
9 "github.com/gosuda/portal/v2/types"
11 - "github.com/gosuda/portal/v2/utils"
10 )
11
14 -func mustSignedRelayDescriptor(t *testing.T, ownerPrivateKey, relayID, relayURL string) types.RelayDescriptor {
12 +func mustRelayDescriptor(t *testing.T, relayID, relayURL string) types.RelayDescriptor {
13 t.Helper()
14
17 - identity, err := utils.ResolveSecp256k1Identity(ownerPrivateKey)
18 - if err != nil {
19 - t.Fatalf("ResolveSecp256k1Identity() error = %v", err)
20 - }
21 -
15 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 - }, identity.PrivateKey)
16 + desc, err := discovery.NormalizeDescriptor(types.RelayDescriptor{
17 + RelayID: relayID,
18 + Sequence: uint64(now.UnixMilli()),
19 + Version: 1,
20 + IssuedAt: now,
21 + ExpiresAt: now.Add(time.Hour),
22 + APIHTTPSAddr: relayURL,
23 + })
24 if err != nil {
34 - t.Fatalf("SignedDescriptor() error = %v", err)
25 + t.Fatalf("NormalizeDescriptor() error = %v", err)
26 }
27 return desc
28 }
@@ -169,22 +160,16 @@ func TestExposureSetRelayURLsRemovesStaleListener(t *testing.T) {
160 }
161 }
162
172 -func TestExposurePinDiscoveredDescriptorRejectsIdentityChange(t *testing.T) {
163 +func TestExposurePinDiscoveredDescriptorRejectsURLChange(t *testing.T) {
164 exposure := &Exposure{relaySet: discovery.NewRelaySet()}
174 - desc := mustSignedRelayDescriptor(t, strings.Repeat("11", 32), "relay-a", "https://relay-a.example")
165 + desc := mustRelayDescriptor(t, "relay-a", "https://relay-a.example")
166
167 if _, _, _, _, err := exposure.relaySet.ApplyRelayDiscoveryResponse(desc.RelayID, desc.APIHTTPSAddr, types.DiscoveryResponse{ProtocolVersion: types.ProtocolVersion, Self: desc}, time.Now().UTC()); err != nil {
168 t.Fatalf("ApplyRelayDiscoveryResponse() error = %v", err)
169 }
170
180 - changedSigner := mustSignedRelayDescriptor(t, strings.Repeat("12", 32), desc.RelayID, desc.APIHTTPSAddr)
181 - _, _, _, _, err := exposure.relaySet.ApplyRelayDiscoveryResponse(desc.RelayID, desc.APIHTTPSAddr, types.DiscoveryResponse{ProtocolVersion: types.ProtocolVersion, Self: changedSigner}, time.Now().UTC())
182 - if err == nil {
183 - t.Fatal("ApplyRelayDiscoveryResponse() error = nil, want pinned signer mismatch")
184 - }
185 -
186 - changedURL := mustSignedRelayDescriptor(t, strings.Repeat("11", 32), desc.RelayID, "https://relay-b.example")
187 - _, _, _, _, err = exposure.relaySet.ApplyRelayDiscoveryResponse(desc.RelayID, "", types.DiscoveryResponse{ProtocolVersion: types.ProtocolVersion, Self: changedURL}, time.Now().UTC())
171 + changedURL := mustRelayDescriptor(t, desc.RelayID, "https://relay-b.example")
172 + _, _, _, _, err := exposure.relaySet.ApplyRelayDiscoveryResponse(desc.RelayID, "", types.DiscoveryResponse{ProtocolVersion: types.ProtocolVersion, Self: changedURL}, time.Now().UTC())
173 if err == nil {
174 t.Fatal("ApplyRelayDiscoveryResponse() error = nil, want pinned relay url mismatch")
175 }
types/discovery.go
-5
@@ -5,9 +5,6 @@ import "time"
5 type RelayDescriptor struct {
6 RelayID string `json:"relay_id"`
7
8 - OwnerAddress string `json:"owner_address"`
9 - SignerPublicKey string `json:"signer_public_key"`
10 -
8 Sequence uint64 `json:"sequence"`
9 Version uint32 `json:"version"`
10 IssuedAt time.Time `json:"issued_at"`
@@ -24,8 +21,6 @@ type RelayDescriptor struct {
21 SupportsTCP bool `json:"supports_tcp,omitempty"`
22 SupportsUDP bool `json:"supports_udp,omitempty"`
23 SupportsOverlayPeer bool `json:"supports_overlay_peer,omitempty"`
27 -
28 - DescriptorSignature string `json:"descriptor_signature"`
24 }
25
26 type RelayState struct {