refactor codes
Kim committed
Mar 26, 2026 at 21:50 UTC
a95221908d18eb07e7796ed3e9232c8ad161373d
18 files changed
+1407
-1685
cmd/relay-server/main.go
+1
-1
@@ -44,8 +44,8 @@ type relayServerConfig struct {
44
AdminSecretKey string
45
TrustProxyHeaders bool
46
TrustedProxyCIDRs string
47
- KeylessDir string
47
AdminSettingsPath string
48
+ KeylessDir string
49
ACMEDNSProvider string
50
CloudflareToken string
51
AWSAccessKeyID string
portal/api_server.go
+52
-58
@@ -36,9 +36,13 @@ var (
36
)
37
38
func (s *Server) newAPIServer(listener net.Listener, apiMux *http.ServeMux, apiTLS keyless.TLSMaterialConfig) (net.Listener, *http.Server, io.Closer, error) {
39
- keylessSignerHandler, err := newKeylessSignerHandler(apiTLS)
40
- if err != nil {
41
- return nil, nil, nil, err
39
+ var keylessSignerHandler http.Handler
40
+ if len(apiTLS.KeyPEM) > 0 {
41
+ signer, err := keyless.NewSigner(apiTLS.KeyPEM)
42
+ if err != nil {
43
+ return nil, nil, nil, fmt.Errorf("configure api signer: %w", err)
44
+ }
45
+ keylessSignerHandler = signer.Handler()
46
}
47
48
apiServer := &http.Server{
@@ -171,9 +175,9 @@ func (s *Server) discoverySelfDescriptor() (types.RelayDescriptor, error) {
175
ingressAddr = fmt.Sprintf("%s:%d", ingressAddr, s.cfg.SNIPort)
176
}
177
174
- supportsOverlayPeer := strings.TrimSpace(s.cfg.WireGuardPublicKey) != "" &&
175
- strings.TrimSpace(s.cfg.WireGuardEndpoint) != "" &&
176
- strings.TrimSpace(s.cfg.OverlayIPv4) != ""
178
+ supportsOverlayPeer := strings.TrimSpace(s.wgConfig.PublicKey) != "" &&
179
+ strings.TrimSpace(s.wgConfig.Endpoint) != "" &&
180
+ strings.TrimSpace(s.wgConfig.OverlayIPv4) != ""
181
182
descriptor := types.RelayDescriptor{
183
RelayID: s.cfg.PortalURL,
@@ -193,10 +197,10 @@ func (s *Server) discoverySelfDescriptor() (types.RelayDescriptor, error) {
197
StatusState: "healthy",
198
}
199
if supportsOverlayPeer {
196
- descriptor.WireGuardPublicKey = strings.TrimSpace(s.cfg.WireGuardPublicKey)
197
- descriptor.WireGuardEndpoint = strings.TrimSpace(s.cfg.WireGuardEndpoint)
198
- descriptor.OverlayIPv4 = strings.TrimSpace(s.cfg.OverlayIPv4)
199
- descriptor.OverlayCIDRs = append([]string(nil), s.cfg.OverlayCIDRs...)
200
+ descriptor.WireGuardPublicKey = strings.TrimSpace(s.wgConfig.PublicKey)
201
+ descriptor.WireGuardEndpoint = strings.TrimSpace(s.wgConfig.Endpoint)
202
+ descriptor.OverlayIPv4 = strings.TrimSpace(s.wgConfig.OverlayIPv4)
203
+ descriptor.OverlayCIDRs = append([]string(nil), s.wgConfig.OverlayCIDRs...)
204
}
205
return discovery.SignedDescriptor(descriptor, s.ownerIdentity.PrivateKey)
206
}
@@ -375,7 +379,16 @@ func (s *Server) handleConnect(w http.ResponseWriter, r *http.Request) {
379
return
380
}
381
378
- lease, err := s.admitLeaseByID(leaseID, token, false)
382
+ lease, err := s.registry.FindByID(leaseID)
383
+ if err == nil && !s.registry.policy.IsLeaseRoutable(lease.ID) {
384
+ err = errLeaseRejected
385
+ }
386
+ if err == nil && !utils.TokenMatches(lease.ReverseToken, token) {
387
+ err = errUnauthorized
388
+ }
389
+ if err == nil && lease.stream == nil {
390
+ err = errTransportMismatch
391
+ }
392
switch {
393
case errors.Is(err, errLeaseNotFound):
394
utils.WriteAPIError(w, http.StatusNotFound, types.APIErrorCodeLeaseNotFound, err.Error())
@@ -415,17 +428,11 @@ func (s *Server) handleConnect(w http.ResponseWriter, r *http.Request) {
428
return
429
}
430
418
- stream := lease.stream
419
- if stream == nil {
420
- _ = conn.Close()
421
- return
422
- }
423
-
431
remoteAddr := ""
432
if conn.RemoteAddr() != nil {
433
remoteAddr = conn.RemoteAddr().String()
434
}
428
- if err := stream.OfferConn(conn); err != nil {
435
+ if err := lease.stream.OfferConn(conn); err != nil {
436
log.Warn().
437
Err(err).
438
Str("lease_id", lease.ID).
@@ -440,7 +447,7 @@ func (s *Server) handleConnect(w http.ResponseWriter, r *http.Request) {
447
Str("lease_id", lease.ID).
448
Str("lease_name", lease.Name).
449
Str("remote_addr", remoteAddr).
443
- Int("ready", stream.ReadyCount()).
450
+ Int("ready", lease.stream.ReadyCount()).
451
Msg("sdk reverse connected")
452
}
453
@@ -464,7 +471,16 @@ func (s *Server) handleQUICTunnelConn(conn *quic.Conn) {
471
return
472
}
473
467
- lease, err := s.admitLeaseByID(msg.LeaseID, msg.ReverseToken, true)
474
+ lease, err := s.registry.FindByID(msg.LeaseID)
475
+ if err == nil && !s.registry.policy.IsLeaseRoutable(lease.ID) {
476
+ err = errLeaseRejected
477
+ }
478
+ if err == nil && !utils.TokenMatches(lease.ReverseToken, msg.ReverseToken) {
479
+ err = errUnauthorized
480
+ }
481
+ if err == nil && (lease.stream == nil || lease.datagram == nil) {
482
+ err = errTransportMismatch
483
+ }
484
switch {
485
case errors.Is(err, errLeaseNotFound):
486
_ = json.NewEncoder(stream).Encode(types.QUICControlResponse{OK: false, Error: types.APIErrorCodeLeaseNotFound})
@@ -488,13 +504,7 @@ func (s *Server) handleQUICTunnelConn(conn *quic.Conn) {
504
return
505
}
506
491
- dg := lease.datagram
492
- if dg == nil {
493
- _ = json.NewEncoder(stream).Encode(types.QUICControlResponse{OK: false, Error: types.APIErrorCodeTransportMismatch})
494
- _ = conn.CloseWithError(1, "transport mismatch")
495
- return
496
- }
497
- if err := dg.Register(conn); err != nil {
507
+ if err := lease.datagram.Register(conn); err != nil {
508
_ = json.NewEncoder(stream).Encode(types.QUICControlResponse{OK: false, Error: "broker_closed"})
509
_ = conn.CloseWithError(1, "broker closed")
510
return
@@ -536,17 +546,16 @@ func (s *Server) registerLease(req types.RegisterRequest, clientIP string) (type
546
}
547
ownerAddress := strings.TrimSpace(req.OwnerAddress)
548
if ownerAddress != "" {
539
- ownerAddress, err = discovery.NormalizeEVMAddress(ownerAddress)
549
+ ownerAddress, err = utils.NormalizeEVMAddress(ownerAddress)
550
if err != nil {
551
return types.RegisterResponse{}, fmt.Errorf("normalize owner address: %w", err)
552
}
553
}
554
545
- if err := s.requireDatagramPlane(req.UDPEnabled); err != nil {
546
- return types.RegisterResponse{}, err
547
- }
548
-
555
if req.UDPEnabled {
556
+ if s.cfg.UDPPortCount <= 0 || s.group != nil && s.quicTunnel == nil {
557
+ return types.RegisterResponse{}, errFeatureUnavailable
558
+ }
559
if !s.registry.policy.IsUDPEnabled() {
560
return types.RegisterResponse{}, errUDPDisabled
561
}
@@ -612,13 +621,18 @@ func (s *Server) registerLease(req types.RegisterRequest, clientIP string) (type
621
return types.RegisterResponse{}, err
622
}
623
if s.DiscoveryEnabled() {
615
- advertisedURLs, relayURLErr := discovery.RelayAPIURLs(s.discoveryCache.AdvertisedDescriptors())
616
- if relayURLErr != nil {
624
+ advertisedURLs := make([]string, 0, len(s.discoveryCache.AdvertisedDescriptors()))
625
+ for _, descriptor := range s.discoveryCache.AdvertisedDescriptors() {
626
+ if apiURL := strings.TrimSpace(descriptor.APIHTTPSAddr); apiURL != "" {
627
+ advertisedURLs = append(advertisedURLs, apiURL)
628
+ }
629
+ }
630
+ responseBootstraps, err = utils.ExcludeLocalRelayURLs(append(responseBootstraps, advertisedURLs...)...)
631
+ if err != nil {
632
record.Close()
633
_, _ = s.registry.Unregister(record.ID, record.ReverseToken)
619
- return types.RegisterResponse{}, relayURLErr
634
+ return types.RegisterResponse{}, err
635
}
621
- responseBootstraps, err = utils.NormalizeRelayURLs(append(responseBootstraps, advertisedURLs...)...)
636
} else {
637
responseBootstraps, err = utils.NormalizeRelayURLs(append(responseBootstraps, append(s.cfg.Bootstraps, record.Bootstraps...)...)...)
638
}
@@ -666,16 +680,8 @@ func (s *Server) unregisterLease(req types.UnregisterRequest) error {
680
if err != nil {
681
return err
682
}
669
- s.closeLease(record)
670
- return nil
671
-}
672
-
673
-func (s *Server) authorizeLeaseToken(record *leaseRecord, token string) error {
674
- if record == nil {
675
- return errLeaseNotFound
676
- }
677
- if !utils.TokenMatches(record.ReverseToken, token) {
678
- return errUnauthorized
683
+ if record != nil {
684
+ record.Close()
685
}
686
return nil
687
}
@@ -687,15 +693,3 @@ func (s *Server) runAPIServer() error {
693
}
694
return err
695
}
690
-
691
-func newKeylessSignerHandler(apiTLS keyless.TLSMaterialConfig) (http.Handler, error) {
692
- if len(apiTLS.KeyPEM) == 0 {
693
- return nil, nil
694
- }
695
-
696
- signer, err := keyless.NewSigner(apiTLS.KeyPEM)
697
- if err != nil {
698
- return nil, fmt.Errorf("configure api signer: %w", err)
699
- }
700
- return signer.Handler(), nil
701
-}
portal/discovery/descriptor.go
deleted
-299
@@ -1,299 +0,0 @@
1
-package discovery
2
-
3
-import (
4
- "crypto/sha256"
5
- "encoding/base64"
6
- "encoding/hex"
7
- "encoding/json"
8
- "errors"
9
- "fmt"
10
- "net"
11
- "sort"
12
- "strconv"
13
- "strings"
14
- "time"
15
-
16
- "github.com/decred/dcrd/dcrec/secp256k1/v4"
17
- secp256k1ecdsa "github.com/decred/dcrd/dcrec/secp256k1/v4/ecdsa"
18
-
19
- "github.com/gosuda/portal/v2/types"
20
- "github.com/gosuda/portal/v2/utils"
21
-)
22
-
23
-func NormalizeDescriptor(desc types.RelayDescriptor) (types.RelayDescriptor, error) {
24
- desc.RelayID = strings.TrimSpace(desc.RelayID)
25
- desc.OwnerAddress = strings.TrimSpace(desc.OwnerAddress)
26
- desc.SignerPublicKey = strings.ToLower(strings.TrimSpace(desc.SignerPublicKey))
27
- desc.APIHTTPSAddr = strings.TrimSpace(desc.APIHTTPSAddr)
28
- desc.IngressTLSAddr = strings.TrimSpace(desc.IngressTLSAddr)
29
- desc.WireGuardPublicKey = strings.TrimSpace(desc.WireGuardPublicKey)
30
- desc.WireGuardEndpoint = strings.TrimSpace(desc.WireGuardEndpoint)
31
- desc.OverlayIPv4 = strings.TrimSpace(desc.OverlayIPv4)
32
- desc.StatusState = strings.TrimSpace(desc.StatusState)
33
- desc.Region = strings.TrimSpace(desc.Region)
34
- desc.Country = strings.TrimSpace(desc.Country)
35
- desc.DescriptorSignature = strings.ToLower(strings.TrimSpace(desc.DescriptorSignature))
36
- if !desc.IssuedAt.IsZero() {
37
- desc.IssuedAt = desc.IssuedAt.UTC()
38
- }
39
- if !desc.ExpiresAt.IsZero() {
40
- desc.ExpiresAt = desc.ExpiresAt.UTC()
41
- }
42
- if !desc.LastMITMDetectedAt.IsZero() {
43
- desc.LastMITMDetectedAt = desc.LastMITMDetectedAt.UTC()
44
- }
45
-
46
- if desc.APIHTTPSAddr != "" {
47
- normalized, err := utils.NormalizeRelayURL(desc.APIHTTPSAddr)
48
- if err != nil {
49
- return types.RelayDescriptor{}, fmt.Errorf("normalize api https addr: %w", err)
50
- }
51
- desc.APIHTTPSAddr = normalized
52
- if desc.RelayID == "" {
53
- desc.RelayID = normalized
54
- }
55
- }
56
- if desc.OwnerAddress != "" {
57
- address, err := NormalizeEVMAddress(desc.OwnerAddress)
58
- if err != nil {
59
- return types.RelayDescriptor{}, fmt.Errorf("normalize owner address: %w", err)
60
- }
61
- desc.OwnerAddress = address
62
- }
63
- if len(desc.OverlayCIDRs) > 0 {
64
- normalized, err := NormalizeOverlayCIDRs(desc.OverlayCIDRs)
65
- if err != nil {
66
- return types.RelayDescriptor{}, err
67
- }
68
- desc.OverlayCIDRs = normalized
69
- }
70
-
71
- if !desc.SupportsOverlayPeer {
72
- desc.WireGuardPublicKey = ""
73
- desc.WireGuardEndpoint = ""
74
- desc.OverlayIPv4 = ""
75
- desc.OverlayCIDRs = nil
76
- }
77
-
78
- return desc, nil
79
-}
80
-
81
-func CanonicalDescriptorPayload(desc types.RelayDescriptor) ([]byte, error) {
82
- normalized, err := NormalizeDescriptor(desc)
83
- if err != nil {
84
- return nil, err
85
- }
86
- normalized.DescriptorSignature = ""
87
- return json.Marshal(normalized)
88
-}
89
-
90
-func SignDescriptor(desc types.RelayDescriptor, privateKeyHex string) (string, error) {
91
- keyHex := strings.TrimSpace(privateKeyHex)
92
- if keyHex == "" {
93
- return "", errors.New("private key is required")
94
- }
95
- if strings.HasPrefix(strings.ToLower(keyHex), "0x") {
96
- keyHex = keyHex[2:]
97
- }
98
- decoded, err := hex.DecodeString(keyHex)
99
- if err != nil {
100
- return "", errors.New("private key must be hex encoded")
101
- }
102
- if len(decoded) != secp256k1.PrivKeyBytesLen {
103
- return "", fmt.Errorf("private key must be %d bytes", secp256k1.PrivKeyBytesLen)
104
- }
105
-
106
- payload, err := CanonicalDescriptorPayload(desc)
107
- if err != nil {
108
- return "", err
109
- }
110
- hash := sha256.Sum256(payload)
111
- privateKey := secp256k1.PrivKeyFromBytes(decoded)
112
- signature := secp256k1ecdsa.Sign(privateKey, hash[:])
113
- return hex.EncodeToString(signature.Serialize()), nil
114
-}
115
-
116
-func SignedDescriptor(desc types.RelayDescriptor, privateKeyHex string) (types.RelayDescriptor, error) {
117
- normalized, err := NormalizeDescriptor(desc)
118
- if err != nil {
119
- return types.RelayDescriptor{}, err
120
- }
121
- signature, err := SignDescriptor(normalized, privateKeyHex)
122
- if err != nil {
123
- return types.RelayDescriptor{}, err
124
- }
125
- normalized.DescriptorSignature = signature
126
- return normalized, nil
127
-}
128
-
129
-func VerifyDescriptor(desc types.RelayDescriptor) error {
130
- normalized, err := NormalizeDescriptor(desc)
131
- if err != nil {
132
- return err
133
- }
134
- if normalized.DescriptorSignature == "" {
135
- return errors.New("descriptor signature is required")
136
- }
137
- if normalized.SignerPublicKey == "" {
138
- return errors.New("signer public key is required")
139
- }
140
-
141
- pubKeyBytes, err := hex.DecodeString(normalized.SignerPublicKey)
142
- if err != nil {
143
- return errors.New("signer public key must be hex encoded")
144
- }
145
- pubKey, err := secp256k1.ParsePubKey(pubKeyBytes)
146
- if err != nil {
147
- return errors.New("invalid secp256k1 signer public key")
148
- }
149
-
150
- sigBytes, err := hex.DecodeString(normalized.DescriptorSignature)
151
- if err != nil {
152
- return errors.New("descriptor signature must be hex encoded")
153
- }
154
- signature, err := secp256k1ecdsa.ParseDERSignature(sigBytes)
155
- if err != nil {
156
- return fmt.Errorf("parse descriptor signature: %w", err)
157
- }
158
-
159
- payload, err := CanonicalDescriptorPayload(normalized)
160
- if err != nil {
161
- return err
162
- }
163
- hash := sha256.Sum256(payload)
164
- if !signature.Verify(hash[:], pubKey) {
165
- return errors.New("descriptor signature is invalid")
166
- }
167
- return nil
168
-}
169
-
170
-func ValidateDescriptor(desc types.RelayDescriptor, now time.Time) (types.RelayDescriptor, error) {
171
- normalized, err := NormalizeDescriptor(desc)
172
- if err != nil {
173
- return types.RelayDescriptor{}, err
174
- }
175
- if now.IsZero() {
176
- now = time.Now()
177
- }
178
- now = now.UTC()
179
-
180
- switch {
181
- case normalized.RelayID == "":
182
- return types.RelayDescriptor{}, errors.New("relay_id is required")
183
- case normalized.OwnerAddress == "":
184
- return types.RelayDescriptor{}, errors.New("owner_address is required")
185
- case normalized.SignerPublicKey == "":
186
- return types.RelayDescriptor{}, errors.New("signer_public_key is required")
187
- case normalized.APIHTTPSAddr == "":
188
- return types.RelayDescriptor{}, errors.New("api_https_addr is required")
189
- case normalized.Sequence == 0:
190
- return types.RelayDescriptor{}, errors.New("sequence is required")
191
- case normalized.Version == 0:
192
- return types.RelayDescriptor{}, errors.New("version is required")
193
- case normalized.IssuedAt.IsZero():
194
- return types.RelayDescriptor{}, errors.New("issued_at is required")
195
- case normalized.ExpiresAt.IsZero():
196
- return types.RelayDescriptor{}, errors.New("expires_at is required")
197
- case normalized.ExpiresAt.Before(now):
198
- return types.RelayDescriptor{}, errors.New("descriptor expired")
199
- case normalized.IssuedAt.After(normalized.ExpiresAt):
200
- return types.RelayDescriptor{}, errors.New("issued_at must be before expires_at")
201
- }
202
-
203
- derivedOwnerAddress, err := AddressFromCompressedPublicKeyHex(normalized.SignerPublicKey)
204
- if err != nil {
205
- return types.RelayDescriptor{}, err
206
- }
207
- if normalized.OwnerAddress != derivedOwnerAddress {
208
- return types.RelayDescriptor{}, errors.New("owner_address does not match signer_public_key")
209
- }
210
-
211
- if normalized.SupportsOverlayPeer {
212
- if err := ValidateWireGuardPublicKey(normalized.WireGuardPublicKey); err != nil {
213
- return types.RelayDescriptor{}, err
214
- }
215
- if err := ValidateWireGuardEndpoint(normalized.WireGuardEndpoint); err != nil {
216
- return types.RelayDescriptor{}, err
217
- }
218
- if err := ValidateOverlayIPv4(normalized.OverlayIPv4); err != nil {
219
- return types.RelayDescriptor{}, err
220
- }
221
- }
222
-
223
- if err := VerifyDescriptor(normalized); err != nil {
224
- return types.RelayDescriptor{}, err
225
- }
226
- return normalized, nil
227
-}
228
-
229
-func ValidateWireGuardPublicKey(raw string) error {
230
- key := strings.TrimSpace(raw)
231
- if key == "" {
232
- return errors.New("wireguard_public_key is required")
233
- }
234
- decoded, err := base64.StdEncoding.DecodeString(key)
235
- if err != nil {
236
- return errors.New("wireguard_public_key must be base64 encoded")
237
- }
238
- if len(decoded) != 32 {
239
- return errors.New("wireguard_public_key must be 32 bytes")
240
- }
241
- return nil
242
-}
243
-
244
-func ValidateWireGuardEndpoint(raw string) error {
245
- endpoint := strings.TrimSpace(raw)
246
- if endpoint == "" {
247
- return errors.New("wireguard_endpoint is required")
248
- }
249
- host, port, err := net.SplitHostPort(endpoint)
250
- if err != nil {
251
- return errors.New("wireguard_endpoint must be host:port")
252
- }
253
- if strings.TrimSpace(host) == "" {
254
- return errors.New("wireguard_endpoint host is required")
255
- }
256
- portNum, err := strconv.Atoi(port)
257
- if err != nil || portNum <= 0 || portNum > 65535 {
258
- return errors.New("wireguard_endpoint port is invalid")
259
- }
260
- return nil
261
-}
262
-
263
-func ValidateOverlayIPv4(raw string) error {
264
- ipText := strings.TrimSpace(raw)
265
- if ipText == "" {
266
- return errors.New("overlay_ipv4 is required")
267
- }
268
- ip := net.ParseIP(ipText)
269
- if ip == nil || ip.To4() == nil {
270
- return errors.New("overlay_ipv4 must be a valid IPv4 address")
271
- }
272
- return nil
273
-}
274
-
275
-func NormalizeOverlayCIDRs(inputs []string) ([]string, error) {
276
- if len(inputs) == 0 {
277
- return nil, nil
278
- }
279
- seen := make(map[string]struct{}, len(inputs))
280
- out := make([]string, 0, len(inputs))
281
- for _, input := range inputs {
282
- input = strings.TrimSpace(input)
283
- if input == "" {
284
- continue
285
- }
286
- _, network, err := net.ParseCIDR(input)
287
- if err != nil {
288
- return nil, fmt.Errorf("invalid overlay cidr %q", input)
289
- }
290
- normalized := network.String()
291
- if _, ok := seen[normalized]; ok {
292
- continue
293
- }
294
- seen[normalized] = struct{}{}
295
- out = append(out, normalized)
296
- }
297
- sort.Strings(out)
298
- return out, nil
299
-}
portal/discovery/discovery.go
+196
-134
@@ -3,6 +3,7 @@ package discovery
3
import (
4
"context"
5
"crypto/tls"
6
+ "encoding/json"
7
"errors"
8
"fmt"
9
"net/http"
@@ -19,170 +20,183 @@ type Resolver func(context.Context, types.DiscoverRequest) (types.DiscoverRespon
20
21
const defaultRequestTimeout = 15 * time.Second
22
22
-func Discover(ctx context.Context, relayURL string, req types.DiscoverRequest, rootCAPEM []byte) (types.DiscoverResponse, error) {
23
- return discoverPeer(ctx, relayURL, req, rootCAPEM)
24
-}
25
-
26
-func RelayAPIURLs(descriptors []types.RelayDescriptor) ([]string, error) {
27
- urls := make([]string, 0, len(descriptors))
28
- for _, descriptor := range descriptors {
29
- if apiURL := strings.TrimSpace(descriptor.APIHTTPSAddr); apiURL != "" {
30
- urls = append(urls, apiURL)
31
- }
23
+func NormalizeDescriptor(desc types.RelayDescriptor) (types.RelayDescriptor, error) {
24
+ desc.RelayID = strings.TrimSpace(desc.RelayID)
25
+ desc.SignerPublicKey = strings.ToLower(strings.TrimSpace(desc.SignerPublicKey))
26
+ desc.APIHTTPSAddr = strings.TrimSpace(desc.APIHTTPSAddr)
27
+ desc.WireGuardPublicKey = strings.TrimSpace(desc.WireGuardPublicKey)
28
+ desc.WireGuardEndpoint = strings.TrimSpace(desc.WireGuardEndpoint)
29
+ desc.OverlayIPv4 = strings.TrimSpace(desc.OverlayIPv4)
30
+ desc.DescriptorSignature = strings.TrimSpace(desc.DescriptorSignature)
31
+ if !desc.IssuedAt.IsZero() {
32
+ desc.IssuedAt = desc.IssuedAt.UTC()
33
}
33
- if len(urls) == 0 {
34
- return nil, nil
34
+ if !desc.ExpiresAt.IsZero() {
35
+ desc.ExpiresAt = desc.ExpiresAt.UTC()
36
}
36
-
37
- normalized, err := utils.NormalizeRelayURLs(urls...)
38
- if err != nil {
39
- return nil, err
37
+ if !desc.LastMITMDetectedAt.IsZero() {
38
+ desc.LastMITMDetectedAt = desc.LastMITMDetectedAt.UTC()
39
}
41
- return utils.ExcludeLocalRelayURLs(normalized...)
42
-}
40
44
-func ResolvePeerResponse(resp types.DiscoverResponse, now time.Time) (types.RelayDescriptor, []types.RelayDescriptor, error) {
45
- self, err := ValidateDescriptor(resp.Self, now)
46
- if err != nil {
47
- return types.RelayDescriptor{}, nil, fmt.Errorf("validate self descriptor: %w", err)
41
+ if desc.APIHTTPSAddr != "" {
42
+ normalized, err := utils.NormalizeRelayURL(desc.APIHTTPSAddr)
43
+ if err != nil {
44
+ return types.RelayDescriptor{}, fmt.Errorf("normalize api https addr: %w", err)
45
+ }
46
+ desc.APIHTTPSAddr = normalized
47
+ if desc.RelayID == "" {
48
+ desc.RelayID = normalized
49
+ }
50
}
49
-
50
- seen := map[string]struct{}{self.RelayID: {}}
51
- peers := make([]types.RelayDescriptor, 0, len(resp.Peers))
52
- var resolveErr error
53
-
54
- for _, descriptor := range resp.Peers {
55
- verified, err := ValidateDescriptor(descriptor, now)
51
+ if desc.OwnerAddress != "" {
52
+ address, err := utils.NormalizeEVMAddress(desc.OwnerAddress)
53
if err != nil {
57
- resolveErr = errors.Join(resolveErr, fmt.Errorf("validate peer %q: %w", descriptor.RelayID, err))
58
- continue
54
+ return types.RelayDescriptor{}, fmt.Errorf("normalize owner address: %w", err)
55
}
60
- if _, ok := seen[verified.RelayID]; ok {
61
- continue
56
+ desc.OwnerAddress = address
57
+ }
58
+ if len(desc.OverlayCIDRs) > 0 {
59
+ normalized, err := utils.NormalizeOverlayCIDRs(desc.OverlayCIDRs)
60
+ if err != nil {
61
+ return types.RelayDescriptor{}, err
62
}
63
- seen[verified.RelayID] = struct{}{}
64
- peers = append(peers, verified)
63
+ desc.OverlayCIDRs = normalized
64
}
66
-
67
- return self, peers, resolveErr
65
+ if !desc.SupportsOverlayPeer {
66
+ desc.WireGuardPublicKey = ""
67
+ desc.WireGuardEndpoint = ""
68
+ desc.OverlayIPv4 = ""
69
+ desc.OverlayCIDRs = nil
70
+ }
71
+ return desc, nil
72
}
73
70
-func DiscoverBootstraps(ctx context.Context, peers []string, req types.DiscoverRequest, rootCAPEM []byte) ([]string, error) {
71
- peers, err := utils.ExcludeLocalRelayURLs(peers...)
74
+func SignDescriptor(desc types.RelayDescriptor, privateKeyHex string) (string, error) {
75
+ normalized, err := NormalizeDescriptor(desc)
76
if err != nil {
73
- return nil, err
77
+ return "", err
78
}
75
- if len(peers) == 0 {
76
- return nil, nil
79
+ normalized.DescriptorSignature = ""
80
+ payload, err := json.Marshal(normalized)
81
+ if err != nil {
82
+ return "", err
83
}
84
+ return utils.SignSHA256Secp256k1DER(payload, privateKeyHex)
85
+}
86
79
- req, err = normalizeRequest(req)
87
+func SignedDescriptor(desc types.RelayDescriptor, privateKeyHex string) (types.RelayDescriptor, error) {
88
+ normalized, err := NormalizeDescriptor(desc)
89
if err != nil {
81
- return nil, err
90
+ return types.RelayDescriptor{}, err
91
}
83
-
84
- bootstraps := append([]string(nil), peers...)
85
- var discoverErr error
86
- discovered := false
87
-
88
- for _, peer := range peers {
89
- resp, err := Discover(ctx, peer, req, rootCAPEM)
90
- if err != nil {
91
- discoverErr = errors.Join(discoverErr, fmt.Errorf("discover %q: %w", peer, err))
92
- continue
93
- }
94
-
95
- self, advertised, resolveErr := ResolvePeerResponse(resp, time.Now().UTC())
96
- if strings.TrimSpace(self.RelayID) == "" {
97
- discoverErr = errors.Join(discoverErr, fmt.Errorf("resolve %q self descriptor: %w", peer, resolveErr))
98
- continue
99
- }
100
- if resolveErr != nil {
101
- discoverErr = errors.Join(discoverErr, fmt.Errorf("resolve %q descriptors: %w", peer, resolveErr))
102
- }
103
-
104
- descriptors := append([]types.RelayDescriptor{self}, advertised...)
105
- discoveredBootstraps, err := RelayAPIURLs(descriptors)
106
- if err != nil {
107
- discoverErr = errors.Join(discoverErr, fmt.Errorf("extract %q relay urls: %w", peer, err))
108
- continue
109
- }
110
- bootstraps, err = utils.MergeRelayURLs(bootstraps, nil, discoveredBootstraps)
111
- if err != nil {
112
- discoverErr = errors.Join(discoverErr, fmt.Errorf("merge %q bootstraps: %w", peer, err))
113
- continue
114
- }
115
- discovered = true
92
+ signature, err := SignDescriptor(normalized, privateKeyHex)
93
+ if err != nil {
94
+ return types.RelayDescriptor{}, err
95
}
96
+ normalized.DescriptorSignature = signature
97
+ return normalized, nil
98
+}
99
118
- if !discovered {
119
- return bootstraps, discoverErr
100
+func ValidateDescriptor(desc types.RelayDescriptor, now time.Time) (types.RelayDescriptor, error) {
101
+ normalized, err := NormalizeDescriptor(desc)
102
+ if err != nil {
103
+ return types.RelayDescriptor{}, err
104
}
121
- return bootstraps, discoverErr
122
-}
105
+ if now.IsZero() {
106
+ now = time.Now()
107
+ }
108
+ now = now.UTC()
109
124
-func ServeHTTP(w http.ResponseWriter, r *http.Request, resolver Resolver) {
125
- if r.Method != http.MethodGet {
126
- utils.WriteAPIError(w, http.StatusMethodNotAllowed, types.APIErrorCodeMethodNotAllowed, "method not allowed")
127
- return
110
+ switch {
111
+ case normalized.RelayID == "":
112
+ return types.RelayDescriptor{}, errors.New("relay_id is required")
113
+ case normalized.APIHTTPSAddr == "":
114
+ return types.RelayDescriptor{}, errors.New("api_https_addr is required")
115
+ case normalized.Sequence == 0:
116
+ return types.RelayDescriptor{}, errors.New("sequence is required")
117
+ case normalized.Version == 0:
118
+ return types.RelayDescriptor{}, errors.New("version is required")
119
+ case normalized.IssuedAt.IsZero():
120
+ return types.RelayDescriptor{}, errors.New("issued_at is required")
121
+ case normalized.ExpiresAt.IsZero():
122
+ return types.RelayDescriptor{}, errors.New("expires_at is required")
123
+ case normalized.ExpiresAt.Before(now):
124
+ return types.RelayDescriptor{}, errors.New("descriptor expired")
125
+ case normalized.IssuedAt.After(normalized.ExpiresAt):
126
+ return types.RelayDescriptor{}, errors.New("issued_at must be before expires_at")
127
}
128
130
- req, err := normalizeRequest(types.DiscoverRequest{
131
- RootHost: r.URL.Query().Get("root_host"),
132
- Name: r.URL.Query().Get("name"),
133
- })
129
+ derivedOwnerAddress, err := utils.AddressFromCompressedPublicKeyHex(normalized.SignerPublicKey)
130
if err != nil {
135
- utils.WriteAPIError(w, http.StatusBadRequest, types.APIErrorCodeInvalidRequest, err.Error())
136
- return
131
+ return types.RelayDescriptor{}, err
132
}
138
- if resolver == nil {
139
- utils.WriteAPIError(w, http.StatusInternalServerError, types.APIErrorCodeInternal, "discovery resolver is not configured")
140
- return
133
+ if normalized.OwnerAddress != derivedOwnerAddress {
134
+ return types.RelayDescriptor{}, errors.New("owner_address does not match signer_public_key")
135
+ }
136
+ if normalized.SupportsOverlayPeer {
137
+ if err := utils.ValidateWireGuardPublicKey(normalized.WireGuardPublicKey); err != nil {
138
+ return types.RelayDescriptor{}, err
139
+ }
140
+ if err := utils.ValidateWireGuardEndpoint(normalized.WireGuardEndpoint); err != nil {
141
+ return types.RelayDescriptor{}, err
142
+ }
143
+ if err := utils.ValidateOverlayIPv4(normalized.OverlayIPv4); err != nil {
144
+ return types.RelayDescriptor{}, err
145
+ }
146
}
147
143
- resp, err := resolver(r.Context(), req)
148
+ signature := normalized.DescriptorSignature
149
+ normalized.DescriptorSignature = ""
150
+ payload, err := json.Marshal(normalized)
151
if err != nil {
145
- utils.WriteAPIError(w, http.StatusInternalServerError, types.APIErrorCodeInternal, err.Error())
146
- return
152
+ return types.RelayDescriptor{}, err
153
}
148
- utils.WriteAPIData(w, http.StatusOK, resp)
154
+ if err := utils.VerifySHA256Secp256k1DER(payload, normalized.SignerPublicKey, signature); err != nil {
155
+ return types.RelayDescriptor{}, err
156
+ }
157
+ normalized.DescriptorSignature = signature
158
+ return normalized, nil
159
}
160
151
-func normalizeRequest(req types.DiscoverRequest) (types.DiscoverRequest, error) {
152
- req.RootHost = utils.NormalizeHostname(req.RootHost)
153
- req.Name = strings.TrimSpace(req.Name)
154
- if req.Name == "" {
155
- return req, nil
156
- }
157
- if req.RootHost == "" {
158
- return types.DiscoverRequest{}, errors.New("root host is required when name is set")
159
- }
160
- name, err := utils.NormalizeDNSLabel(req.Name)
161
+func ValidateResponse(resp types.DiscoverResponse, now time.Time) (types.RelayDescriptor, []types.RelayDescriptor, error) {
162
+ self, err := ValidateDescriptor(resp.Self, now)
163
if err != nil {
162
- return types.DiscoverRequest{}, err
164
+ return types.RelayDescriptor{}, nil, err
165
}
164
- req.Name = name
165
- return req, nil
166
-}
166
168
-func discoverPeer(ctx context.Context, relayURL string, req types.DiscoverRequest, rootCAPEM []byte) (types.DiscoverResponse, error) {
169
- relayURL, err := utils.NormalizeRelayURL(relayURL)
170
- if err != nil {
171
- return types.DiscoverResponse{}, err
167
+ seen := map[string]struct{}{self.RelayID: {}}
168
+ peers := make([]types.RelayDescriptor, 0, len(resp.Peers))
169
+ var validateErr error
170
+ for _, descriptor := range resp.Peers {
171
+ verified, err := ValidateDescriptor(descriptor, now)
172
+ if err != nil {
173
+ validateErr = errors.Join(validateErr, fmt.Errorf("validate peer %q: %w", descriptor.RelayID, err))
174
+ continue
175
+ }
176
+ if _, ok := seen[verified.RelayID]; ok {
177
+ continue
178
+ }
179
+ seen[verified.RelayID] = struct{}{}
180
+ peers = append(peers, verified)
181
}
182
+ return self, peers, validateErr
183
+}
184
174
- baseURL, err := url.Parse(relayURL)
175
- if err != nil {
176
- return types.DiscoverResponse{}, fmt.Errorf("parse relay url: %w", err)
185
+func Discover(ctx context.Context, baseURL string, req types.DiscoverRequest, rootCAPEM []byte, httpClient *http.Client) (types.DiscoverResponse, error) {
186
+ baseURL = strings.TrimSpace(baseURL)
187
+ if baseURL == "" {
188
+ return types.DiscoverResponse{}, errors.New("discovery base url is required")
189
}
190
179
- rootCAs, err := keyless.RelayRootCAs(ctx, relayURL, baseURL.Hostname(), rootCAPEM)
191
+ parsedBaseURL, err := url.Parse(baseURL)
192
if err != nil {
181
- return types.DiscoverResponse{}, err
193
+ return types.DiscoverResponse{}, fmt.Errorf("parse discovery base url: %w", err)
194
+ }
195
+ if parsedBaseURL.Host == "" {
196
+ return types.DiscoverResponse{}, errors.New("discovery base url host is required")
197
}
198
184
- ref, _ := url.Parse(types.PathDiscovery)
185
- discoverURL := baseURL.ResolveReference(ref)
199
+ discoverURL := parsedBaseURL.ResolveReference(&url.URL{Path: types.PathDiscovery})
200
query := discoverURL.Query()
201
if req.RootHost != "" {
202
query.Set("root_host", req.RootHost)
@@ -192,17 +206,28 @@ func discoverPeer(ctx context.Context, relayURL string, req types.DiscoverReques
206
}
207
discoverURL.RawQuery = query.Encode()
208
195
- httpClient := &http.Client{
196
- Transport: &http.Transport{
197
- TLSClientConfig: &tls.Config{
198
- MinVersion: tls.VersionTLS12,
199
- ServerName: baseURL.Hostname(),
200
- RootCAs: rootCAs,
201
- NextProtos: []string{"http/1.1"},
209
+ client := httpClient
210
+ if client == nil {
211
+ rootCAs, err := keyless.RelayRootCAs(ctx, baseURL, parsedBaseURL.Hostname(), rootCAPEM)
212
+ if err != nil {
213
+ return types.DiscoverResponse{}, err
214
+ }
215
+ client = &http.Client{
216
+ Transport: &http.Transport{
217
+ TLSClientConfig: &tls.Config{
218
+ MinVersion: tls.VersionTLS12,
219
+ ServerName: parsedBaseURL.Hostname(),
220
+ RootCAs: rootCAs,
221
+ NextProtos: []string{"http/1.1"},
222
+ },
223
+ ForceAttemptHTTP2: false,
224
},
203
- ForceAttemptHTTP2: false,
204
- },
205
- Timeout: defaultRequestTimeout,
225
+ }
226
+ }
227
+ if client.Timeout == 0 {
228
+ clone := *client
229
+ clone.Timeout = defaultRequestTimeout
230
+ client = &clone
231
}
232
233
httpReq, err := http.NewRequestWithContext(ctx, http.MethodGet, discoverURL.String(), nil)
@@ -210,7 +235,7 @@ func discoverPeer(ctx context.Context, relayURL string, req types.DiscoverReques
235
return types.DiscoverResponse{}, err
236
}
237
213
- resp, err := httpClient.Do(httpReq)
238
+ resp, err := client.Do(httpReq)
239
if err != nil {
240
return types.DiscoverResponse{}, err
241
}
@@ -229,3 +254,40 @@ func discoverPeer(ctx context.Context, relayURL string, req types.DiscoverReques
254
}
255
return envelope.Data, nil
256
}
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/identity.go
deleted
-152
@@ -1,152 +0,0 @@
1
-package discovery
2
-
3
-import (
4
- "encoding/hex"
5
- "errors"
6
- "fmt"
7
- "strings"
8
-
9
- "github.com/decred/dcrd/dcrec/secp256k1/v4"
10
- "golang.org/x/crypto/sha3"
11
-)
12
-
13
-type Identity struct {
14
- Generated bool `json:"generated,omitempty"`
15
- Address string `json:"address"`
16
- PublicKey string `json:"public_key"`
17
- PrivateKey string `json:"private_key"`
18
-}
19
-
20
-func AddressFromCompressedPublicKeyHex(rawPublicKey string) (string, error) {
21
- publicKeyHex := strings.TrimSpace(rawPublicKey)
22
- if publicKeyHex == "" {
23
- return "", errors.New("public key is required")
24
- }
25
- if strings.HasPrefix(strings.ToLower(publicKeyHex), "0x") {
26
- publicKeyHex = publicKeyHex[2:]
27
- }
28
-
29
- decoded, err := hex.DecodeString(publicKeyHex)
30
- if err != nil {
31
- return "", errors.New("public key must be hex encoded")
32
- }
33
-
34
- publicKey, err := secp256k1.ParsePubKey(decoded)
35
- if err != nil {
36
- return "", errors.New("invalid secp256k1 public key")
37
- }
38
-
39
- uncompressed := publicKey.SerializeUncompressed()
40
- if len(uncompressed) != 65 || uncompressed[0] != 0x04 {
41
- return "", errors.New("invalid uncompressed secp256k1 public key")
42
- }
43
-
44
- hasher := sha3.NewLegacyKeccak256()
45
- _, _ = hasher.Write(uncompressed[1:])
46
- hash := hasher.Sum(nil)
47
-
48
- return NormalizeEVMAddress("0x" + hex.EncodeToString(hash[len(hash)-20:]))
49
-}
50
-
51
-func NormalizeEVMAddress(raw string) (string, error) {
52
- trimmed := strings.TrimSpace(raw)
53
- if trimmed == "" {
54
- return "", errors.New("address is required")
55
- }
56
- if !strings.HasPrefix(strings.ToLower(trimmed), "0x") {
57
- return "", errors.New("address must start with 0x")
58
- }
59
-
60
- hexPart := trimmed[2:]
61
- if len(hexPart) != 40 {
62
- return "", errors.New("address must be 20 bytes")
63
- }
64
- if _, err := hex.DecodeString(hexPart); err != nil {
65
- return "", errors.New("address must be hex encoded")
66
- }
67
-
68
- lowerHex := strings.ToLower(hexPart)
69
- hasher := sha3.NewLegacyKeccak256()
70
- _, _ = hasher.Write([]byte(lowerHex))
71
- hash := hasher.Sum(nil)
72
-
73
- var builder strings.Builder
74
- builder.Grow(len(lowerHex))
75
- for idx, ch := range lowerHex {
76
- if ch >= '0' && ch <= '9' {
77
- builder.WriteRune(ch)
78
- continue
79
- }
80
-
81
- nibble := hash[idx/2]
82
- if idx%2 == 0 {
83
- nibble >>= 4
84
- } else {
85
- nibble &= 0x0f
86
- }
87
- if nibble > 7 {
88
- builder.WriteRune(ch - ('a' - 'A'))
89
- continue
90
- }
91
- builder.WriteRune(ch)
92
- }
93
-
94
- checksummed := builder.String()
95
- if hexPart != lowerHex && hexPart != strings.ToUpper(hexPart) && hexPart != checksummed {
96
- return "", errors.New("address checksum is invalid")
97
- }
98
- return "0x" + checksummed, nil
99
-}
100
-
101
-func ResolveIdentity(rawPrivateKey string) (Identity, error) {
102
- privateKeyHex := strings.TrimSpace(rawPrivateKey)
103
- generated := false
104
- if privateKeyHex == "" {
105
- privateKey, err := secp256k1.GeneratePrivateKey()
106
- if err != nil {
107
- return Identity{}, fmt.Errorf("generate secp256k1 private key: %w", err)
108
- }
109
- privateKeyHex = hex.EncodeToString(privateKey.Serialize())
110
- generated = true
111
- }
112
- if strings.HasPrefix(strings.ToLower(privateKeyHex), "0x") {
113
- privateKeyHex = privateKeyHex[2:]
114
- }
115
-
116
- decoded, err := hex.DecodeString(privateKeyHex)
117
- if err != nil {
118
- return Identity{}, errors.New("secp256k1 private key must be hex encoded")
119
- }
120
- if len(decoded) != secp256k1.PrivKeyBytesLen {
121
- return Identity{}, fmt.Errorf("secp256k1 private key must be %d bytes", secp256k1.PrivKeyBytesLen)
122
- }
123
-
124
- isZero := true
125
- for _, b := range decoded {
126
- if b != 0 {
127
- isZero = false
128
- break
129
- }
130
- }
131
- if isZero {
132
- return Identity{}, errors.New("secp256k1 private key must not be zero")
133
- }
134
-
135
- privateKey := secp256k1.PrivKeyFromBytes(decoded)
136
- if privateKey == nil {
137
- return Identity{}, errors.New("invalid secp256k1 private key")
138
- }
139
-
140
- publicKeyHex := hex.EncodeToString(privateKey.PubKey().SerializeCompressed())
141
- address, err := AddressFromCompressedPublicKeyHex(publicKeyHex)
142
- if err != nil {
143
- return Identity{}, err
144
- }
145
-
146
- return Identity{
147
- Generated: generated,
148
- Address: address,
149
- PublicKey: publicKeyHex,
150
- PrivateKey: privateKeyHex,
151
- }, nil
152
-}
portal/discovery/store.go
+4
-31
@@ -25,9 +25,9 @@ type Cache struct {
25
peers map[string]peerRecord
26
}
27
28
-func (s *Cache) Lookup(relayID string) (types.PeerState, bool) {
28
+func (s *Cache) Lookup(relayID string) (types.PeerState, bool, bool) {
29
if strings.TrimSpace(relayID) == "" {
30
- return types.PeerState{}, false
30
+ return types.PeerState{}, false, false
31
}
32
33
s.mu.RLock()
@@ -35,9 +35,9 @@ func (s *Cache) Lookup(relayID string) (types.PeerState, bool) {
35
36
record, ok := s.peers[relayID]
37
if !ok {
38
- return types.PeerState{}, false
38
+ return types.PeerState{}, false, false
39
}
40
- return record.state, true
40
+ return record.state, strings.TrimSpace(record.pinnedSignerPublicKey) != "", true
41
}
42
43
func NewCache() *Cache {
@@ -126,21 +126,6 @@ func (s *Cache) Snapshot() map[string]types.PeerState {
126
return out
127
}
128
129
-func (s *Cache) SeedURL(relayID string) string {
130
- if strings.TrimSpace(relayID) == "" {
131
- return ""
132
- }
133
-
134
- s.mu.RLock()
135
- defer s.mu.RUnlock()
136
-
137
- record, ok := s.peers[relayID]
138
- if !ok {
139
- return ""
140
- }
141
- return record.seedURL
142
-}
143
-
129
func (s *Cache) KnownDescriptors() []types.RelayDescriptor {
130
s.mu.RLock()
131
defer s.mu.RUnlock()
@@ -175,18 +160,6 @@ func (s *Cache) AdvertisedDescriptors() []types.RelayDescriptor {
160
return out
161
}
162
178
-func (s *Cache) HasPinnedIdentity(relayID string) bool {
179
- if strings.TrimSpace(relayID) == "" {
180
- return false
181
- }
182
-
183
- s.mu.RLock()
184
- defer s.mu.RUnlock()
185
-
186
- record, ok := s.peers[relayID]
187
- return ok && strings.TrimSpace(record.pinnedSignerPublicKey) != ""
188
-}
189
-
163
func (s *Cache) PinIdentity(relayID, seedURL string, desc types.RelayDescriptor) error {
164
relayID = strings.TrimSpace(relayID)
165
if relayID == "" {
portal/discovery/store_test.go
deleted
-90
@@ -1,90 +0,0 @@
1
-package discovery
2
-
3
-import (
4
- "strings"
5
- "testing"
6
- "time"
7
-
8
- "github.com/gosuda/portal/v2/types"
9
-)
10
-
11
-func signedRelayDescriptor(t *testing.T, privateKey, relayURL string) types.RelayDescriptor {
12
- t.Helper()
13
-
14
- identity, err := ResolveIdentity(privateKey)
15
- if err != nil {
16
- t.Fatalf("ResolveIdentity() error = %v", err)
17
- }
18
-
19
- now := time.Now().UTC()
20
- desc, err := SignedDescriptor(types.RelayDescriptor{
21
- RelayID: relayURL,
22
- OwnerAddress: identity.Address,
23
- SignerPublicKey: identity.PublicKey,
24
- Sequence: uint64(now.UnixMilli()),
25
- Version: 1,
26
- IssuedAt: now,
27
- ExpiresAt: now.Add(time.Hour),
28
- APIHTTPSAddr: relayURL,
29
- SupportsTCP: true,
30
- StatusState: "healthy",
31
- }, identity.PrivateKey)
32
- if err != nil {
33
- t.Fatalf("SignedDescriptor() error = %v", err)
34
- }
35
- return desc
36
-}
37
-
38
-func TestCacheRecordVerifiedReportsDescriptorChanges(t *testing.T) {
39
- t.Parallel()
40
-
41
- cache := NewCache()
42
- if _, err := cache.UpsertSeedURLs([]string{"https://relay-a.example.com"}); err != nil {
43
- t.Fatalf("UpsertSeedURLs() error = %v", err)
44
- }
45
-
46
- desc := signedRelayDescriptor(t, strings.Repeat("11", 32), "https://relay-a.example.com")
47
- if err := cache.PinIdentity(desc.RelayID, desc.APIHTTPSAddr, desc); err != nil {
48
- t.Fatalf("PinIdentity() error = %v", err)
49
- }
50
-
51
- added, changed, err := cache.RecordVerified(desc, true)
52
- if err != nil {
53
- t.Fatalf("RecordVerified() error = %v", err)
54
- }
55
- if added || !changed {
56
- t.Fatalf("RecordVerified() = added:%v changed:%v, want false true", added, changed)
57
- }
58
-
59
- updated := desc
60
- updated.StatusState = "degraded"
61
- updated.DescriptorSignature, err = SignDescriptor(updated, strings.Repeat("11", 32))
62
- if err != nil {
63
- t.Fatalf("SignDescriptor() error = %v", err)
64
- }
65
-
66
- added, changed, err = cache.RecordVerified(updated, true)
67
- if err != nil {
68
- t.Fatalf("RecordVerified() second error = %v", err)
69
- }
70
- if added || !changed {
71
- t.Fatalf("RecordVerified() second = added:%v changed:%v, want false true", added, changed)
72
- }
73
-}
74
-
75
-func TestCacheKnownDescriptorsIncludeExpiredForRehydration(t *testing.T) {
76
- t.Parallel()
77
-
78
- cache := NewCache()
79
- if _, err := cache.UpsertSeedURLs([]string{"https://relay-a.example.com"}); err != nil {
80
- t.Fatalf("UpsertSeedURLs() error = %v", err)
81
- }
82
- if !cache.Expire("https://relay-a.example.com") {
83
- t.Fatal("Expire() = false, want true")
84
- }
85
-
86
- known := cache.KnownDescriptors()
87
- if len(known) != 1 || known[0].RelayID != "https://relay-a.example.com" {
88
- t.Fatalf("KnownDescriptors() = %+v, want expired relay retained for rehydration", known)
89
- }
90
-}
portal/server.go
+234
-457
@@ -8,13 +8,13 @@ import (
8
"io"
9
"net"
10
"net/http"
11
- "sort"
11
"strings"
12
"sync"
13
"time"
14
15
"github.com/gosuda/keyless_tls/relay/l4"
16
"github.com/quic-go/quic-go"
17
+ "github.com/rs/zerolog"
18
"github.com/rs/zerolog/log"
19
"golang.org/x/sync/errgroup"
20
@@ -71,17 +71,16 @@ type Server struct {
71
sniListener net.Listener
72
apiListener net.Listener
73
apiServer *http.Server
74
- wgPeerListener net.Listener
75
- wgPeerServer *http.Server
74
apiTLSClose io.Closer
75
acmeManager *acme.Manager
76
quicTunnel *quic.Listener
79
- wgRuntime *wireguard.Runtime
77
+ overlay *wireguard.Overlay
78
cancel context.CancelFunc
79
group *errgroup.Group
80
registry *leaseRegistry
81
ports *transport.PortAllocator
84
- ownerIdentity discovery.Identity
82
+ ownerIdentity utils.Secp256k1Identity
83
+ wgConfig wireguard.Config
84
cfg ServerConfig
85
rootHost string
86
trustedProxyCIDRs []*net.IPNet
@@ -114,49 +113,16 @@ func NewServer(cfg ServerConfig) (*Server, error) {
113
return nil, fmt.Errorf("normalize bootstraps: %w", err)
114
}
115
cfg.Bootstraps = bootstraps
117
- wireGuardConfigured := strings.TrimSpace(cfg.WireGuardPrivateKey) != "" ||
118
- strings.TrimSpace(cfg.WireGuardPublicKey) != "" ||
119
- strings.TrimSpace(cfg.WireGuardEndpoint) != "" ||
120
- strings.TrimSpace(cfg.OverlayIPv4) != "" ||
121
- len(cfg.OverlayCIDRs) > 0
122
- if wireGuardConfigured {
123
- if strings.TrimSpace(cfg.WireGuardPrivateKey) == "" {
124
- return nil, errors.New("wireguard private key is required when relay overlay is enabled")
125
- }
126
- cfg.WireGuardPrivateKey, err = utils.NormalizeWireGuardPrivateKey(cfg.WireGuardPrivateKey)
127
- if err != nil {
128
- return nil, fmt.Errorf("normalize wireguard private key: %w", err)
129
- }
130
- derivedPublicKey, err := utils.WireGuardPublicKeyFromPrivate(cfg.WireGuardPrivateKey)
131
- if err != nil {
132
- return nil, fmt.Errorf("derive wireguard public key: %w", err)
133
- }
134
- if configuredPublicKey := strings.TrimSpace(cfg.WireGuardPublicKey); configuredPublicKey != "" && configuredPublicKey != derivedPublicKey {
135
- return nil, errors.New("wireguard public key does not match private key")
136
- }
137
- cfg.WireGuardPublicKey = derivedPublicKey
138
- cfg.DiscoveryPort = utils.IntOrDefault(cfg.DiscoveryPort, wireguard.DefaultListenPort)
139
- if len(cfg.OverlayCIDRs) > 0 {
140
- cfg.OverlayCIDRs, err = discovery.NormalizeOverlayCIDRs(cfg.OverlayCIDRs)
141
- if err != nil {
142
- return nil, fmt.Errorf("normalize overlay cidrs: %w", err)
143
- }
144
- }
145
- if strings.TrimSpace(cfg.WireGuardEndpoint) == "" {
146
- cfg.WireGuardEndpoint = net.JoinHostPort(rootHost, fmt.Sprintf("%d", cfg.DiscoveryPort))
147
- }
148
- if strings.TrimSpace(cfg.OverlayIPv4) == "" {
149
- cfg.OverlayIPv4, err = utils.DeriveWireGuardOverlayIPv4(cfg.WireGuardPublicKey)
150
- if err != nil {
151
- return nil, fmt.Errorf("derive overlay ipv4: %w", err)
152
- }
153
- }
154
- if err := discovery.ValidateWireGuardEndpoint(cfg.WireGuardEndpoint); err != nil {
155
- return nil, err
156
- }
157
- if err := discovery.ValidateOverlayIPv4(cfg.OverlayIPv4); err != nil {
158
- return nil, err
159
- }
116
+ wgConfig, err := wireguard.NormalizeConfig(rootHost, wireguard.Config{
117
+ PrivateKey: cfg.WireGuardPrivateKey,
118
+ PublicKey: cfg.WireGuardPublicKey,
119
+ Endpoint: cfg.WireGuardEndpoint,
120
+ OverlayIPv4: cfg.OverlayIPv4,
121
+ OverlayCIDRs: cfg.OverlayCIDRs,
122
+ ListenPort: cfg.DiscoveryPort,
123
+ })
124
+ if err != nil {
125
+ return nil, err
126
}
127
128
portMin, portMax := 0, 0
@@ -166,7 +132,7 @@ func NewServer(cfg ServerConfig) (*Server, error) {
132
}
133
134
ownerPrivateKey := strings.TrimSpace(cfg.OwnerPrivateKey)
169
- ownerIdentity, err := discovery.ResolveIdentity(ownerPrivateKey)
135
+ ownerIdentity, err := utils.ResolveSecp256k1Identity(ownerPrivateKey)
136
if err != nil {
137
if ownerPrivateKey == "" {
138
return nil, fmt.Errorf("generate relay owner private key: %w", err)
@@ -192,12 +158,15 @@ func NewServer(cfg ServerConfig) (*Server, error) {
158
registry: registry,
159
ports: ports,
160
ownerIdentity: ownerIdentity,
161
+ wgConfig: wgConfig,
162
trustedProxyCIDRs: trustedProxyCIDRs,
163
}
164
165
// Tear down all lease resources when leases expire via TTL janitor.
166
registry.onExpired = func(record *leaseRecord) {
200
- s.closeLease(record)
167
+ if record != nil {
168
+ record.Close()
169
+ }
170
}
171
172
if cfg.DiscoveryEnabled {
@@ -254,27 +223,57 @@ func (s *Server) Start(ctx context.Context, apiMux *http.ServeMux) error {
223
s.cancel = cancel
224
s.group = group
225
257
- if s.wireGuardPeerPlaneEnabled() {
258
- if err := s.startWireGuardPeerPlane(); err != nil {
226
+ if s.wgConfig.PrivateKey != "" {
227
+ var snapshot map[string]types.PeerState
228
+ if s.discoveryCache != nil {
229
+ snapshot = s.discoveryCache.Snapshot()
230
+ }
231
+ peerMux := http.NewServeMux()
232
+ peerMux.HandleFunc(types.PathRoot, s.handleRoot)
233
+ peerMux.HandleFunc(types.PathHealthz, s.handleHealthz)
234
+ peerMux.HandleFunc(types.PathDiscovery, func(w http.ResponseWriter, r *http.Request) {
235
+ if !s.DiscoveryEnabled() {
236
+ http.NotFound(w, r)
237
+ return
238
+ }
239
+ discovery.ServeHTTP(w, r, s.discover)
240
+ })
241
+ overlay, err := wireguard.NewOverlay(s.wgConfig, peerMux)
242
+ if err != nil {
243
+ acmeManager.Stop()
244
+ _ = apiServer.Close()
245
+ _ = apiCloser.Close()
246
+ _ = sniListener.Close()
247
+ cancel()
248
+ return fmt.Errorf("start wireguard overlay: %w", err)
249
+ }
250
+ if err := overlay.Sync(s.cfg.PortalURL, snapshot); err != nil {
251
acmeManager.Stop()
252
_ = apiServer.Close()
253
_ = apiCloser.Close()
254
_ = sniListener.Close()
255
+ _ = overlay.Shutdown(context.Background())
256
cancel()
264
- return fmt.Errorf("start wireguard peer plane: %w", err)
257
+ return fmt.Errorf("sync wireguard peers: %w", err)
258
}
259
+ s.overlay = overlay
260
}
261
262
group.Go(s.runAPIServer)
269
- if s.wgPeerServer != nil {
270
- group.Go(s.runWireGuardPeerAPIServer)
263
+ if s.overlay != nil {
264
+ group.Go(s.overlay.Serve)
265
}
266
group.Go(func() error { return s.runSNIListener(groupCtx) })
267
group.Go(func() error { return s.registry.RunJanitor(groupCtx, 5*time.Second) })
268
if s.DiscoveryEnabled() {
269
group.Go(func() error { return s.runDiscoveryLoop(groupCtx) })
270
}
277
- group.Go(func() error { return s.watchContext(groupCtx) })
271
+ group.Go(func() error {
272
+ <-groupCtx.Done()
273
+ shutdownCtx, cancel := context.WithTimeout(context.Background(), 5*time.Second)
274
+ defer cancel()
275
+ return s.Shutdown(shutdownCtx)
276
+ })
277
s.acmeManager.Start(serverCtx)
278
279
if s.cfg.UDPPortCount > 0 {
@@ -301,7 +300,9 @@ func (s *Server) Shutdown(ctx context.Context) error {
300
}
301
302
for _, lease := range s.registry.CloseAll() {
304
- s.closeLease(lease)
303
+ if lease != nil {
304
+ lease.Close()
305
+ }
306
}
307
308
if s.quicTunnel != nil {
@@ -317,13 +318,8 @@ func (s *Server) Shutdown(ctx context.Context) error {
318
shutdownErr = err
319
}
320
}
320
- if s.wgPeerServer != nil {
321
- if err := s.wgPeerServer.Shutdown(ctx); err != nil && shutdownErr == nil && !errors.Is(err, http.ErrServerClosed) {
322
- shutdownErr = err
323
- }
324
- }
325
- if s.wgRuntime != nil {
326
- _ = s.wgRuntime.Close()
321
+ if err := s.overlay.Shutdown(ctx); err != nil && shutdownErr == nil {
322
+ shutdownErr = err
323
}
324
if s.apiTLSClose != nil {
325
_ = s.apiTLSClose.Close()
@@ -374,16 +370,6 @@ func (s *Server) PortalURL() string {
370
return s.cfg.PortalURL
371
}
372
377
-func (s *Server) wireGuardPeerPlaneEnabled() bool {
378
- if s == nil {
379
- return false
380
- }
381
- return strings.TrimSpace(s.cfg.WireGuardPrivateKey) != "" &&
382
- strings.TrimSpace(s.cfg.WireGuardPublicKey) != "" &&
383
- strings.TrimSpace(s.cfg.WireGuardEndpoint) != "" &&
384
- strings.TrimSpace(s.cfg.OverlayIPv4) != ""
385
-}
386
-
373
func (s *Server) OwnerAddress() string {
374
if s == nil {
375
return ""
@@ -425,68 +411,6 @@ func (s *Server) LeaseSnapshotByHostname(hostname string) (types.Lease, bool) {
411
return s.registry.Snapshot(record), true
412
}
413
428
-func (s *Server) startWireGuardPeerPlane() error {
429
- runtime, err := wireguard.NewRuntime(wireguard.RuntimeConfig{
430
- PrivateKey: s.cfg.WireGuardPrivateKey,
431
- Endpoint: s.cfg.WireGuardEndpoint,
432
- OverlayIPv4: s.cfg.OverlayIPv4,
433
- })
434
- if err != nil {
435
- return err
436
- }
437
-
438
- listener, err := runtime.ListenTCP(wireguard.DefaultPeerAPIHTTPPort)
439
- if err != nil {
440
- _ = runtime.Close()
441
- return fmt.Errorf("listen peer api: %w", err)
442
- }
443
-
444
- server := &http.Server{
445
- Handler: s.peerAPIHandler(),
446
- ReadHeaderTimeout: 10 * time.Second,
447
- }
448
-
449
- s.wgRuntime = runtime
450
- s.wgPeerListener = listener
451
- s.wgPeerServer = server
452
-
453
- if err := s.syncWireGuardPeers(); err != nil {
454
- _ = server.Close()
455
- _ = runtime.Close()
456
- s.wgRuntime = nil
457
- s.wgPeerListener = nil
458
- s.wgPeerServer = nil
459
- return fmt.Errorf("seed wireguard peers: %w", err)
460
- }
461
- return nil
462
-}
463
-
464
-func (s *Server) runWireGuardPeerAPIServer() error {
465
- if s == nil || s.wgPeerServer == nil || s.wgPeerListener == nil {
466
- return nil
467
- }
468
-
469
- err := s.wgPeerServer.Serve(s.wgPeerListener)
470
- if errors.Is(err, http.ErrServerClosed) || errors.Is(err, net.ErrClosed) {
471
- return nil
472
- }
473
- return err
474
-}
475
-
476
-func (s *Server) peerAPIHandler() http.Handler {
477
- mux := http.NewServeMux()
478
- mux.HandleFunc(types.PathRoot, s.handleRoot)
479
- mux.HandleFunc(types.PathHealthz, s.handleHealthz)
480
- mux.HandleFunc(types.PathDiscovery, func(w http.ResponseWriter, r *http.Request) {
481
- if !s.DiscoveryEnabled() {
482
- http.NotFound(w, r)
483
- return
484
- }
485
- discovery.ServeHTTP(w, r, s.discover)
486
- })
487
- return mux
488
-}
489
-
414
func (s *Server) prepareAPITLS(ctx context.Context) (keyless.TLSMaterialConfig, *acme.Manager, error) {
415
acmeCfg := s.cfg.ACME
416
if baseDomain := utils.NormalizeHostname(acmeCfg.BaseDomain); baseDomain != "" && baseDomain != s.rootHost {
@@ -509,22 +433,16 @@ func (s *Server) prepareAPITLS(ctx context.Context) (keyless.TLSMaterialConfig,
433
CertPEM: certPEM,
434
KeyPEM: keyPEM,
435
}
512
- if err := validateAPITLS(apiTLS); err != nil {
513
- manager.Stop()
514
- return keyless.TLSMaterialConfig{}, nil, err
515
- }
516
-
517
- return apiTLS, manager, nil
518
-}
519
-
520
-func validateAPITLS(apiTLS keyless.TLSMaterialConfig) error {
436
if len(apiTLS.CertPEM) == 0 {
522
- return errors.New("api tls certificate is required")
437
+ manager.Stop()
438
+ return keyless.TLSMaterialConfig{}, nil, errors.New("api tls certificate is required")
439
}
440
if len(apiTLS.KeyPEM) == 0 && apiTLS.Keyless == nil {
525
- return errors.New("api tls key or keyless signer is required")
441
+ manager.Stop()
442
+ return keyless.TLSMaterialConfig{}, nil, errors.New("api tls key or keyless signer is required")
443
}
527
- return nil
444
+
445
+ return apiTLS, manager, nil
446
}
447
448
func (s *Server) runSNIListener(ctx context.Context) error {
@@ -532,132 +450,61 @@ func (s *Server) runSNIListener(ctx context.Context) error {
450
conn, err := s.sniListener.Accept()
451
switch {
452
case err == nil:
535
- go s.handleSNIConn(ctx, conn)
536
- case errors.Is(err, net.ErrClosed):
537
- return nil
538
- default:
539
- return err
540
- }
541
- }
542
-}
543
-
544
-func (s *Server) handleSNIConn(ctx context.Context, conn net.Conn) {
545
- clientHello, wrappedConn, err := l4.InspectClientHello(conn, s.cfg.ClientHelloTimeout)
546
- if err != nil {
547
- if wrappedConn != nil {
548
- _ = wrappedConn.Close()
549
- } else {
550
- _ = conn.Close()
551
- }
552
- return
553
- }
554
-
555
- serverName := utils.NormalizeHostname(clientHello.ServerName)
556
- if serverName == "" {
557
- _ = wrappedConn.Close()
558
- return
559
- }
560
-
561
- if serverName == s.rootHost {
562
- s.bridgeToAPI(ctx, wrappedConn)
563
- return
564
- }
565
-
566
- stream, err := s.resolveStream(serverName)
567
- if err != nil {
568
- _ = wrappedConn.Close()
569
- return
570
- }
571
-
572
- claimCtx, cancel := context.WithTimeout(ctx, s.cfg.ClaimTimeout)
573
- defer cancel()
574
-
575
- session, err := stream.Claim(claimCtx)
576
- if err != nil {
577
- _ = wrappedConn.Close()
578
- return
579
- }
580
-
581
- BridgeConns(wrappedConn, session)
582
-}
453
+ go func(conn net.Conn) {
454
+ clientHello, wrappedConn, err := l4.InspectClientHello(conn, s.cfg.ClientHelloTimeout)
455
+ if err != nil {
456
+ if wrappedConn != nil {
457
+ _ = wrappedConn.Close()
458
+ } else {
459
+ _ = conn.Close()
460
+ }
461
+ return
462
+ }
463
584
-func (s *Server) bridgeToAPI(ctx context.Context, conn net.Conn) {
585
- if s.apiListener == nil {
586
- _ = conn.Close()
587
- return
588
- }
589
- dialer := &net.Dialer{Timeout: 5 * time.Second}
590
- upstream, err := dialer.DialContext(ctx, "tcp", utils.HostPortOrLoopback(s.apiListener.Addr().String()))
591
- if err != nil {
592
- _ = conn.Close()
593
- return
594
- }
595
- BridgeConns(conn, upstream)
596
-}
464
+ serverName := utils.NormalizeHostname(clientHello.ServerName)
465
+ if serverName == "" {
466
+ _ = wrappedConn.Close()
467
+ return
468
+ }
469
598
-func (s *Server) lookupRoutableLease(serverName string) (*leaseRecord, error) {
599
- record, ok := s.registry.Lookup(serverName)
600
- if !ok || record == nil {
601
- return nil, errors.New("no route")
602
- }
603
- if time.Now().After(record.ExpiresAt) {
604
- return nil, errors.New("lease expired")
605
- }
606
- if !s.registry.policy.IsLeaseRoutable(record.ID) {
607
- return nil, errors.New("not routable")
608
- }
609
- return record, nil
610
-}
470
+ if serverName == s.rootHost {
471
+ if s.apiListener == nil {
472
+ _ = wrappedConn.Close()
473
+ return
474
+ }
475
+ dialer := &net.Dialer{Timeout: 5 * time.Second}
476
+ upstream, err := dialer.DialContext(ctx, "tcp", utils.HostPortOrLoopback(s.apiListener.Addr().String()))
477
+ if err != nil {
478
+ _ = wrappedConn.Close()
479
+ return
480
+ }
481
+ BridgeConns(wrappedConn, upstream)
482
+ return
483
+ }
484
612
-func (s *Server) resolveStream(serverName string) (*transport.RelayStream, error) {
613
- record, err := s.lookupRoutableLease(serverName)
614
- if err != nil {
615
- return nil, err
616
- }
617
- if record.stream == nil {
618
- return nil, errors.New("transport mismatch")
619
- }
620
- return record.stream, nil
621
-}
485
+ record, ok := s.registry.Lookup(serverName)
486
+ if !ok || record == nil || time.Now().After(record.ExpiresAt) || !s.registry.policy.IsLeaseRoutable(record.ID) || record.stream == nil {
487
+ _ = wrappedConn.Close()
488
+ return
489
+ }
490
623
-func (s *Server) datagramPlaneReady() bool {
624
- if s == nil || s.cfg.UDPPortCount <= 0 {
625
- return false
626
- }
627
- if s.group == nil {
628
- return true
629
- }
630
- return s.quicTunnel != nil
631
-}
491
+ claimCtx, cancel := context.WithTimeout(ctx, s.cfg.ClaimTimeout)
492
+ defer cancel()
493
633
-func (s *Server) requireDatagramPlane(udpEnabled bool) error {
634
- if !udpEnabled {
635
- return nil
636
- }
637
- if s.datagramPlaneReady() {
638
- return nil
639
- }
640
- return errFeatureUnavailable
641
-}
494
+ session, err := record.stream.Claim(claimCtx)
495
+ if err != nil {
496
+ _ = wrappedConn.Close()
497
+ return
498
+ }
499
643
-func (s *Server) admitLeaseByID(leaseID, token string, requireDatagram bool) (*leaseRecord, error) {
644
- record, err := s.registry.FindByID(leaseID)
645
- if err != nil {
646
- return nil, err
647
- }
648
- if !s.registry.policy.IsLeaseRoutable(record.ID) {
649
- return nil, errLeaseRejected
650
- }
651
- if err := s.authorizeLeaseToken(record, token); err != nil {
652
- return nil, err
653
- }
654
- if record.stream == nil {
655
- return nil, errTransportMismatch
656
- }
657
- if requireDatagram && record.datagram == nil {
658
- return nil, errTransportMismatch
500
+ BridgeConns(wrappedConn, session)
501
+ }(conn)
502
+ case errors.Is(err, net.ErrClosed):
503
+ return nil
504
+ default:
505
+ return err
506
+ }
507
}
660
- return record, nil
508
}
509
510
func (s *Server) startQUICTunnelListener(apiTLS keyless.TLSMaterialConfig) error {
@@ -708,208 +555,138 @@ func (s *Server) runQUICTunnelListener(listener *quic.Listener) error {
555
}
556
}
557
711
-func (s *Server) closeLease(record *leaseRecord) {
712
- if record == nil {
713
- return
714
- }
715
- record.Close()
716
-}
717
-
718
-func (s *Server) watchContext(ctx context.Context) error {
719
- <-ctx.Done()
720
- shutdownCtx, cancel := context.WithTimeout(context.Background(), 5*time.Second)
721
- defer cancel()
722
- return s.Shutdown(shutdownCtx)
723
-}
724
-
725
-func (s *Server) desiredWireGuardPeers() []types.DesiredPeer {
726
- if s.wgRuntime == nil {
727
- return nil
728
- }
729
-
730
- snapshot := s.discoveryCache.Snapshot()
731
- peers := make([]types.DesiredPeer, 0, len(snapshot))
732
- for _, state := range snapshot {
733
- if state.State != types.PeerStateVerified && state.State != types.PeerStateAdvertised {
734
- continue
735
- }
736
- desc := state.Descriptor
737
- if desc.RelayID == s.cfg.PortalURL || !desc.SupportsOverlayPeer {
738
- continue
739
- }
740
- if strings.TrimSpace(desc.WireGuardPublicKey) == "" || strings.TrimSpace(desc.WireGuardEndpoint) == "" || strings.TrimSpace(desc.OverlayIPv4) == "" {
741
- continue
742
- }
743
-
744
- allowedIPs := []string{desc.OverlayIPv4 + "/32"}
745
- allowedIPs = append(allowedIPs, desc.OverlayCIDRs...)
746
- peers = append(peers, types.DesiredPeer{
747
- RelayID: desc.RelayID,
748
- WireGuardPublicKey: desc.WireGuardPublicKey,
749
- WireGuardEndpoint: desc.WireGuardEndpoint,
750
- AllowedIPs: allowedIPs,
751
- })
752
- }
753
- sort.Slice(peers, func(i, j int) bool {
754
- return peers[i].RelayID < peers[j].RelayID
755
- })
756
- return peers
757
-}
758
-
759
-func (s *Server) syncWireGuardPeers() error {
760
- if s.wgRuntime == nil {
761
- return nil
762
- }
763
- return s.wgRuntime.ApplyPeers(s.desiredWireGuardPeers())
764
-}
765
-
766
-func (s *Server) discoverRelay(ctx context.Context, peer types.RelayDescriptor) (types.DiscoverResponse, error) {
767
- if peer.SupportsOverlayPeer && s.discoveryCache.HasPinnedIdentity(peer.RelayID) {
768
- state, ok := s.discoveryCache.Lookup(peer.RelayID)
769
- if ok && state.State != types.PeerStateExpired && s.wgRuntime != nil {
770
- if strings.TrimSpace(peer.OverlayIPv4) == "" {
771
- return types.DiscoverResponse{}, errors.New("relay peer is missing overlay ipv4")
772
- }
773
- return s.wgRuntime.Discover(ctx, peer.OverlayIPv4, wireguard.DefaultPeerAPIHTTPPort, types.DiscoverRequest{})
774
- }
775
- }
776
-
777
- seedURL := strings.TrimSpace(s.discoveryCache.SeedURL(peer.RelayID))
778
- if seedURL == "" {
779
- seedURL = strings.TrimSpace(peer.APIHTTPSAddr)
780
- }
781
- if seedURL == "" {
782
- return types.DiscoverResponse{}, errors.New("relay peer is missing seed url")
783
- }
784
- return discovery.Discover(ctx, seedURL, types.DiscoverRequest{}, nil)
785
-}
786
-
558
func (s *Server) runDiscoveryLoop(ctx context.Context) error {
559
ticker := time.NewTicker(defaultDiscoveryInterval)
560
defer ticker.Stop()
561
562
for {
563
peers := s.discoveryCache.KnownDescriptors()
793
- if len(peers) > 0 {
794
- for _, peer := range peers {
795
- resp, err := s.discoverRelay(ctx, peer)
796
- switch {
797
- case err == nil:
798
- selfDescriptor, peerDescriptors, resolveErr := discovery.ResolvePeerResponse(resp, time.Now().UTC())
799
- if strings.TrimSpace(selfDescriptor.RelayID) == "" {
800
- s.discoveryCache.RecordFailure(peer.RelayID)
801
- log.Warn().
802
- Err(resolveErr).
803
- Str("peer", peer.APIHTTPSAddr).
804
- Msg("discovery response missing valid self descriptor")
805
- continue
806
- }
807
- if resolveErr != nil {
808
- log.Warn().
809
- Err(resolveErr).
810
- Str("peer", peer.APIHTTPSAddr).
811
- Msg("discovery response contained invalid peer descriptors")
564
+ var overlayClient *http.Client
565
+ if s.overlay != nil {
566
+ overlayClient = s.overlay.Client()
567
+ }
568
+ for _, peer := range peers {
569
+ discoverURL, discoverClient := peer.APIHTTPSAddr, (*http.Client)(nil)
570
+ if state, pinned, _ := s.discoveryCache.Lookup(peer.RelayID); peer.SupportsOverlayPeer && overlayClient != nil && pinned && state.State != types.PeerStateExpired {
571
+ if peer.OverlayIPv4 == "" {
572
+ err := errors.New("relay peer is missing overlay ipv4")
573
+ s.discoveryCache.RecordFailure(peer.RelayID)
574
+ log.Warn().
575
+ Err(err).
576
+ Str("peer", peer.APIHTTPSAddr).
577
+ Msg("discover peer failed")
578
+ continue
579
+ }
580
+ discoverURL = "http://" + net.JoinHostPort(peer.OverlayIPv4, fmt.Sprintf("%d", wireguard.DefaultPeerAPIHTTPPort))
581
+ discoverClient = overlayClient
582
+ }
583
+ resp, err := discovery.Discover(ctx, discoverURL, types.DiscoverRequest{}, nil, discoverClient)
584
+ if err != nil {
585
+ if ctx.Err() != nil {
586
+ return nil
587
+ }
588
+ s.discoveryCache.RecordFailure(peer.RelayID)
589
+ expireReason := ""
590
+ consecutiveFailures := 0
591
+ state, pinned, _ := s.discoveryCache.Lookup(peer.RelayID)
592
+ if pinned &&
593
+ state.State != types.PeerStateExpired &&
594
+ peer.SupportsOverlayPeer &&
595
+ state.ConsecutiveFailures >= defaultWGRecoveryFailures {
596
+ if removed := s.discoveryCache.Expire(peer.RelayID); removed {
597
+ expireReason = "recovery"
598
+ consecutiveFailures = state.ConsecutiveFailures
599
}
813
- if seedURL := strings.TrimSpace(s.discoveryCache.SeedURL(peer.RelayID)); seedURL != "" {
814
- if err := s.discoveryCache.PinIdentity(peer.RelayID, seedURL, selfDescriptor); err != nil {
815
- s.discoveryCache.RecordFailure(peer.RelayID)
816
- log.Warn().
817
- Err(err).
818
- Str("peer", peer.APIHTTPSAddr).
819
- Msg("discovery peer identity pin failed")
820
- continue
821
- }
600
+ }
601
+ var apiErr *types.APIRequestError
602
+ if expireReason == "" && errors.As(err, &apiErr) &&
603
+ (apiErr.StatusCode == http.StatusForbidden ||
604
+ apiErr.StatusCode == http.StatusNotFound ||
605
+ apiErr.StatusCode == http.StatusGone) {
606
+ if removed := s.discoveryCache.Expire(peer.RelayID); removed {
607
+ expireReason = "status"
608
}
609
+ }
610
+ if expireReason != "" {
611
+ err = errors.Join(err, s.overlay.Sync(s.cfg.PortalURL, s.discoveryCache.Snapshot()))
612
+ }
613
824
- peerSetChanged := false
825
- added, changed, err := s.discoveryCache.RecordVerified(selfDescriptor, true)
826
- if err != nil {
827
- s.discoveryCache.RecordFailure(selfDescriptor.RelayID)
828
- log.Warn().
829
- Err(err).
830
- Str("peer", peer.APIHTTPSAddr).
831
- Msg("record self discovery peer failed")
832
- continue
833
- }
834
- peerSetChanged = peerSetChanged || changed
835
- addedHints := make([]string, 0, len(peerDescriptors))
836
- for _, peerDescriptor := range peerDescriptors {
837
- hintAdded, hintChanged, err := s.discoveryCache.RecordVerified(peerDescriptor, false)
838
- if err != nil {
839
- log.Warn().
840
- Err(err).
841
- Str("peer", peerDescriptor.RelayID).
842
- Msg("record hinted discovery peer failed")
843
- continue
844
- }
845
- peerSetChanged = peerSetChanged || hintChanged
846
- if hintAdded || hintChanged {
847
- addedHints = append(addedHints, peerDescriptor.APIHTTPSAddr)
848
- }
849
- }
850
- if peerSetChanged {
851
- if err := s.syncWireGuardPeers(); err != nil {
852
- log.Warn().
853
- Err(err).
854
- Str("peer", peer.APIHTTPSAddr).
855
- Msg("sync wireguard peers failed")
856
- }
614
+ event := log.Warn().
615
+ Err(err).
616
+ Str("peer", peer.APIHTTPSAddr)
617
+ if expireReason != "" {
618
+ event = event.
619
+ Bool("expired", true).
620
+ Str("reason", expireReason)
621
+ if consecutiveFailures > 0 {
622
+ event = event.Int("consecutive_failures", consecutiveFailures)
623
}
624
+ }
625
+ event.Msg("discover peer failed")
626
+ continue
627
+ }
628
859
- if added || changed || len(addedHints) > 0 {
860
- log.Info().
861
- Str("peer", peer.APIHTTPSAddr).
862
- Bool("discoverable", selfDescriptor.SupportsOverlayPeer).
863
- Int("hint_count", len(addedHints)).
864
- Int("known_count", len(s.discoveryCache.KnownDescriptors())).
865
- Int("advertised_count", len(s.discoveryCache.AdvertisedDescriptors())).
866
- Strs("added_hints", addedHints).
867
- Msg("discovery peer state updated")
868
- }
869
- case ctx.Err() != nil:
870
- return nil
871
- default:
872
- s.discoveryCache.RecordFailure(peer.RelayID)
873
- if state, ok := s.discoveryCache.Lookup(peer.RelayID); ok &&
874
- state.State != types.PeerStateExpired &&
875
- peer.SupportsOverlayPeer &&
876
- s.discoveryCache.HasPinnedIdentity(peer.RelayID) &&
877
- state.ConsecutiveFailures >= defaultWGRecoveryFailures {
878
- if removed := s.discoveryCache.Expire(peer.RelayID); removed {
879
- if err := s.syncWireGuardPeers(); err != nil {
880
- log.Warn().
881
- Err(err).
882
- Str("peer", peer.APIHTTPSAddr).
883
- Msg("sync wireguard peers failed")
884
- }
885
- log.Warn().
886
- Int("consecutive_failures", state.ConsecutiveFailures).
887
- Str("peer", peer.APIHTTPSAddr).
888
- Msg("wireguard discovery failed repeatedly, forcing seed re-hydration")
889
- }
890
- }
891
- var apiErr *types.APIRequestError
892
- if errors.As(err, &apiErr) &&
893
- (apiErr.StatusCode == http.StatusForbidden ||
894
- apiErr.StatusCode == http.StatusNotFound ||
895
- apiErr.StatusCode == http.StatusGone) {
896
- if removed := s.discoveryCache.Expire(peer.RelayID); removed {
897
- if err := s.syncWireGuardPeers(); err != nil {
898
- log.Warn().
899
- Err(err).
900
- Str("peer", peer.APIHTTPSAddr).
901
- Msg("sync wireguard peers failed")
902
- }
903
- log.Info().
904
- Str("peer", peer.APIHTTPSAddr).
905
- Msg("discovery peer removed from advertised set")
906
- }
907
- }
629
+ now := time.Now().UTC()
630
+ selfDescriptor, peerDescriptors, warnErr := discovery.ValidateResponse(resp, now)
631
+ if selfDescriptor.RelayID == "" {
632
+ err = errors.Join(warnErr, errors.New("discover response is missing self descriptor"))
633
+ }
634
+ if err == nil {
635
+ err = s.discoveryCache.PinIdentity(peer.RelayID, peer.APIHTTPSAddr, selfDescriptor)
636
+ }
637
+ added, changed := false, false
638
+ if err == nil {
639
+ added, changed, err = s.discoveryCache.RecordVerified(selfDescriptor, true)
640
+ }
641
+ if err != nil {
642
+ s.discoveryCache.RecordFailure(peer.RelayID)
643
+ log.Warn().
644
+ Err(err).
645
+ Str("peer", peer.APIHTTPSAddr).
646
+ Msg("discover peer failed")
647
+ continue
648
+ }
649
909
- log.Warn().
910
- Err(err).
911
- Str("peer", peer.APIHTTPSAddr).
912
- Msg("discover peer failed")
650
+ peerSetChanged := changed
651
+ addedHintCount := 0
652
+ for _, peerDescriptor := range peerDescriptors {
653
+ hintAdded, hintChanged, err := s.discoveryCache.RecordVerified(peerDescriptor, false)
654
+ if err != nil {
655
+ warnErr = errors.Join(warnErr, fmt.Errorf("record hint %q: %w", peerDescriptor.RelayID, err))
656
+ continue
657
+ }
658
+ peerSetChanged = peerSetChanged || hintChanged
659
+ if hintAdded || hintChanged {
660
+ addedHintCount++
661
+ }
662
+ }
663
+ if peerSetChanged {
664
+ if err := s.overlay.Sync(s.cfg.PortalURL, s.discoveryCache.Snapshot()); err != nil {
665
+ warnErr = errors.Join(warnErr, err)
666
+ }
667
+ }
668
+
669
+ updated := added || changed || addedHintCount > 0
670
+ if updated || warnErr != nil {
671
+ var event *zerolog.Event
672
+ if warnErr != nil {
673
+ event = log.Warn().
674
+ Err(warnErr)
675
+ } else {
676
+ event = log.Info()
677
+ }
678
+ event.
679
+ Str("peer", peer.APIHTTPSAddr).
680
+ Bool("discoverable", selfDescriptor.SupportsOverlayPeer).
681
+ Int("known_count", len(s.discoveryCache.KnownDescriptors())).
682
+ Int("advertised_count", len(s.discoveryCache.AdvertisedDescriptors()))
683
+ if addedHintCount > 0 {
684
+ event.Int("added_hint_count", addedHintCount)
685
+ }
686
+ if updated {
687
+ event.Msg("discovery peer updated")
688
+ } else {
689
+ event.Msg("discover peer completed with warnings")
690
}
691
}
692
}
portal/server_test.go
+69
-35
@@ -20,9 +20,9 @@ import (
20
func mustSignedRelayDescriptor(t *testing.T, ownerPrivateKey, relayURL string) types.RelayDescriptor {
21
t.Helper()
22
23
- identity, err := discovery.ResolveIdentity(ownerPrivateKey)
23
+ identity, err := utils.ResolveSecp256k1Identity(ownerPrivateKey)
24
if err != nil {
25
- t.Fatalf("ResolveIdentity() error = %v", err)
25
+ t.Fatalf("ResolveSecp256k1Identity() error = %v", err)
26
}
27
28
now := time.Now().UTC()
@@ -44,16 +44,6 @@ func mustSignedRelayDescriptor(t *testing.T, ownerPrivateKey, relayURL string) t
44
return desc
45
}
46
47
-func mustRelayAPIURLs(t *testing.T, descriptors []types.RelayDescriptor) []string {
48
- t.Helper()
49
-
50
- urls, err := discovery.RelayAPIURLs(descriptors)
51
- if err != nil {
52
- t.Fatalf("RelayAPIURLs() error = %v", err)
53
- }
54
- return urls
55
-}
56
-
47
func TestServerStartInitializesLocalACMEAndSigner(t *testing.T) {
48
t.Parallel()
49
@@ -152,31 +142,31 @@ func TestNewServerDerivesWireGuardConfigFromPrivateKey(t *testing.T) {
142
t.Fatalf("NewServer() error = %v", err)
143
}
144
155
- if server.cfg.WireGuardPrivateKey == "" {
145
+ if server.wgConfig.PrivateKey == "" {
146
t.Fatal("WireGuardPrivateKey = empty, want normalized key")
147
}
158
- if server.cfg.WireGuardPublicKey == "" {
148
+ if server.wgConfig.PublicKey == "" {
149
t.Fatal("WireGuardPublicKey = empty, want derived key")
150
}
161
- if server.cfg.WireGuardEndpoint != net.JoinHostPort("portal.example.com", "41011") {
162
- t.Fatalf("WireGuardEndpoint = %q, want %q", server.cfg.WireGuardEndpoint, net.JoinHostPort("portal.example.com", "41011"))
151
+ if server.wgConfig.Endpoint != net.JoinHostPort("portal.example.com", "41011") {
152
+ t.Fatalf("WireGuardEndpoint = %q, want %q", server.wgConfig.Endpoint, net.JoinHostPort("portal.example.com", "41011"))
153
}
164
- if server.cfg.OverlayIPv4 == "" {
154
+ if server.wgConfig.OverlayIPv4 == "" {
155
t.Fatal("OverlayIPv4 = empty, want derived overlay address")
156
}
167
- if err := discovery.ValidateWireGuardEndpoint(server.cfg.WireGuardEndpoint); err != nil {
157
+ if err := utils.ValidateWireGuardEndpoint(server.wgConfig.Endpoint); err != nil {
158
t.Fatalf("ValidateWireGuardEndpoint() error = %v", err)
159
}
170
- if err := discovery.ValidateOverlayIPv4(server.cfg.OverlayIPv4); err != nil {
160
+ if err := utils.ValidateOverlayIPv4(server.wgConfig.OverlayIPv4); err != nil {
161
t.Fatalf("ValidateOverlayIPv4() error = %v", err)
162
}
163
174
- wantOverlay, err := utils.DeriveWireGuardOverlayIPv4(server.cfg.WireGuardPublicKey)
164
+ wantOverlay, err := utils.DeriveWireGuardOverlayIPv4(server.wgConfig.PublicKey)
165
if err != nil {
166
t.Fatalf("DeriveWireGuardOverlayIPv4() error = %v", err)
167
}
178
- if server.cfg.OverlayIPv4 != wantOverlay {
179
- t.Fatalf("OverlayIPv4 = %q, want %q", server.cfg.OverlayIPv4, wantOverlay)
168
+ if server.wgConfig.OverlayIPv4 != wantOverlay {
169
+ t.Fatalf("OverlayIPv4 = %q, want %q", server.wgConfig.OverlayIPv4, wantOverlay)
170
}
171
}
172
@@ -190,8 +180,8 @@ func TestNewServerIgnoresDiscoveryPortWithoutWireGuardKey(t *testing.T) {
180
if err != nil {
181
t.Fatalf("NewServer() error = %v", err)
182
}
193
- if server.cfg.WireGuardEndpoint != "" {
194
- t.Fatalf("WireGuardEndpoint = %q, want empty without wireguard key", server.cfg.WireGuardEndpoint)
183
+ if server.wgConfig.Endpoint != "" {
184
+ t.Fatalf("WireGuardEndpoint = %q, want empty without wireguard key", server.wgConfig.Endpoint)
185
}
186
}
187
@@ -254,7 +244,7 @@ func TestRegisterLeaseBuildsUDPEnabledRuntime(t *testing.T) {
244
}
245
t.Cleanup(func() {
246
if record, ok := server.registry.Get(resp.LeaseID); ok {
257
- server.closeLease(record)
247
+ record.Close()
248
}
249
})
250
@@ -280,9 +270,9 @@ func TestServerStartServesOptionalDiscoveryRoutes(t *testing.T) {
270
t.Parallel()
271
272
ownerPrivateKey := strings.Repeat("11", 32)
283
- ownerIdentity, err := discovery.ResolveIdentity(ownerPrivateKey)
273
+ ownerIdentity, err := utils.ResolveSecp256k1Identity(ownerPrivateKey)
274
if err != nil {
285
- t.Fatalf("ResolveIdentity() error = %v", err)
275
+ t.Fatalf("ResolveSecp256k1Identity() error = %v", err)
276
}
277
278
server, err := NewServer(ServerConfig{
@@ -416,8 +406,22 @@ func TestServerUpsertDiscoverySeedURLsSkipsLocalRelayHosts(t *testing.T) {
406
if !reflect.DeepEqual(added, []string{"https://relay-a.example.com"}) {
407
t.Fatalf("UpsertSeedURLs() added = %v, want [%q]", added, "https://relay-a.example.com")
408
}
419
- if !reflect.DeepEqual(mustRelayAPIURLs(t, server.discoveryCache.KnownDescriptors()), []string{"https://bootstrap.example.com", "https://relay-a.example.com"}) {
420
- t.Fatalf("KnownDescriptors() = %v, want [%q %q]", mustRelayAPIURLs(t, server.discoveryCache.KnownDescriptors()), "https://bootstrap.example.com", "https://relay-a.example.com")
409
+ knownRelayURLs, err := utils.ExcludeLocalRelayURLs("https://bootstrap.example.com", "https://relay-a.example.com")
410
+ if err != nil {
411
+ t.Fatalf("ExcludeLocalRelayURLs() error = %v", err)
412
+ }
413
+ knownURLs := make([]string, 0, len(server.discoveryCache.KnownDescriptors()))
414
+ for _, descriptor := range server.discoveryCache.KnownDescriptors() {
415
+ if apiURL := strings.TrimSpace(descriptor.APIHTTPSAddr); apiURL != "" {
416
+ knownURLs = append(knownURLs, apiURL)
417
+ }
418
+ }
419
+ knownURLs, err = utils.ExcludeLocalRelayURLs(knownURLs...)
420
+ if err != nil {
421
+ t.Fatalf("ExcludeLocalRelayURLs() known error = %v", err)
422
+ }
423
+ if !reflect.DeepEqual(knownURLs, knownRelayURLs) {
424
+ t.Fatalf("KnownDescriptors() = %v, want [%q %q]", knownURLs, "https://bootstrap.example.com", "https://relay-a.example.com")
425
}
426
if len(server.discoveryCache.AdvertisedDescriptors()) != 0 {
427
t.Fatalf("AdvertisedDescriptors() = %v, want empty before direct confirmation", server.discoveryCache.AdvertisedDescriptors())
@@ -462,11 +466,31 @@ func TestServerRecordVerifiedDiscoveryPeerRequiresDirectConfirmation(t *testing.
466
if !changed {
467
t.Fatal("RecordVerified() hinted changed = false, want true")
468
}
465
- if !reflect.DeepEqual(mustRelayAPIURLs(t, server.discoveryCache.KnownDescriptors()), []string{"https://bootstrap.example.com", "https://relay-a.example.com"}) {
466
- t.Fatalf("KnownDescriptors() = %v, want [%q %q]", mustRelayAPIURLs(t, server.discoveryCache.KnownDescriptors()), "https://bootstrap.example.com", "https://relay-a.example.com")
469
+ knownURLs := make([]string, 0, len(server.discoveryCache.KnownDescriptors()))
470
+ for _, descriptor := range server.discoveryCache.KnownDescriptors() {
471
+ if apiURL := strings.TrimSpace(descriptor.APIHTTPSAddr); apiURL != "" {
472
+ knownURLs = append(knownURLs, apiURL)
473
+ }
474
+ }
475
+ knownURLs, err = utils.ExcludeLocalRelayURLs(knownURLs...)
476
+ if err != nil {
477
+ t.Fatalf("ExcludeLocalRelayURLs() known error = %v", err)
478
+ }
479
+ if !reflect.DeepEqual(knownURLs, []string{"https://bootstrap.example.com", "https://relay-a.example.com"}) {
480
+ t.Fatalf("KnownDescriptors() = %v, want [%q %q]", knownURLs, "https://bootstrap.example.com", "https://relay-a.example.com")
481
+ }
482
+ advertisedURLs := make([]string, 0, len(server.discoveryCache.AdvertisedDescriptors()))
483
+ for _, descriptor := range server.discoveryCache.AdvertisedDescriptors() {
484
+ if apiURL := strings.TrimSpace(descriptor.APIHTTPSAddr); apiURL != "" {
485
+ advertisedURLs = append(advertisedURLs, apiURL)
486
+ }
487
}
468
- if !reflect.DeepEqual(mustRelayAPIURLs(t, server.discoveryCache.AdvertisedDescriptors()), []string{"https://bootstrap.example.com"}) {
469
- t.Fatalf("AdvertisedDescriptors() = %v, want [%q]", mustRelayAPIURLs(t, server.discoveryCache.AdvertisedDescriptors()), "https://bootstrap.example.com")
488
+ advertisedURLs, err = utils.ExcludeLocalRelayURLs(advertisedURLs...)
489
+ if err != nil {
490
+ t.Fatalf("ExcludeLocalRelayURLs() advertised error = %v", err)
491
+ }
492
+ if !reflect.DeepEqual(advertisedURLs, []string{"https://bootstrap.example.com"}) {
493
+ t.Fatalf("AdvertisedDescriptors() = %v, want [%q]", advertisedURLs, "https://bootstrap.example.com")
494
}
495
496
snapshot := server.discoveryCache.Snapshot()
@@ -487,8 +511,18 @@ func TestServerRecordVerifiedDiscoveryPeerRequiresDirectConfirmation(t *testing.
511
if !changed {
512
t.Fatal("RecordVerified() second changed = false, want true")
513
}
490
- if !reflect.DeepEqual(mustRelayAPIURLs(t, server.discoveryCache.AdvertisedDescriptors()), []string{"https://bootstrap.example.com", "https://relay-a.example.com"}) {
491
- t.Fatalf("AdvertisedDescriptors() = %v, want [%q %q]", mustRelayAPIURLs(t, server.discoveryCache.AdvertisedDescriptors()), "https://bootstrap.example.com", "https://relay-a.example.com")
514
+ advertisedURLs = advertisedURLs[:0]
515
+ for _, descriptor := range server.discoveryCache.AdvertisedDescriptors() {
516
+ if apiURL := strings.TrimSpace(descriptor.APIHTTPSAddr); apiURL != "" {
517
+ advertisedURLs = append(advertisedURLs, apiURL)
518
+ }
519
+ }
520
+ advertisedURLs, err = utils.ExcludeLocalRelayURLs(advertisedURLs...)
521
+ if err != nil {
522
+ t.Fatalf("ExcludeLocalRelayURLs() advertised second error = %v", err)
523
+ }
524
+ if !reflect.DeepEqual(advertisedURLs, []string{"https://bootstrap.example.com", "https://relay-a.example.com"}) {
525
+ t.Fatalf("AdvertisedDescriptors() = %v, want [%q %q]", advertisedURLs, "https://bootstrap.example.com", "https://relay-a.example.com")
526
}
527
}
528
portal/wireguard/overlay.go
new
+191
@@ -0,0 +1,191 @@
1
+package wireguard
2
+
3
+import (
4
+ "context"
5
+ "errors"
6
+ "fmt"
7
+ "net"
8
+ "net/http"
9
+ "sort"
10
+ "strings"
11
+ "time"
12
+
13
+ "github.com/gosuda/portal/v2/types"
14
+ "github.com/gosuda/portal/v2/utils"
15
+)
16
+
17
+type Config struct {
18
+ PrivateKey string
19
+ PublicKey string
20
+ Endpoint string
21
+ OverlayIPv4 string
22
+ OverlayCIDRs []string
23
+ ListenPort int
24
+}
25
+
26
+func NormalizeConfig(rootHost string, cfg Config) (Config, error) {
27
+ configured := strings.TrimSpace(cfg.PrivateKey) != "" ||
28
+ strings.TrimSpace(cfg.PublicKey) != "" ||
29
+ strings.TrimSpace(cfg.Endpoint) != "" ||
30
+ strings.TrimSpace(cfg.OverlayIPv4) != "" ||
31
+ len(cfg.OverlayCIDRs) > 0
32
+ if !configured {
33
+ return cfg, nil
34
+ }
35
+
36
+ if strings.TrimSpace(cfg.PrivateKey) == "" {
37
+ return Config{}, errors.New("wireguard private key is required when relay overlay is enabled")
38
+ }
39
+
40
+ privateKey, err := utils.NormalizeWireGuardPrivateKey(cfg.PrivateKey)
41
+ if err != nil {
42
+ return Config{}, fmt.Errorf("normalize wireguard private key: %w", err)
43
+ }
44
+ publicKey, err := utils.WireGuardPublicKeyFromPrivate(privateKey)
45
+ if err != nil {
46
+ return Config{}, fmt.Errorf("derive wireguard public key: %w", err)
47
+ }
48
+ if configuredPublicKey := strings.TrimSpace(cfg.PublicKey); configuredPublicKey != "" && configuredPublicKey != publicKey {
49
+ return Config{}, errors.New("wireguard public key does not match private key")
50
+ }
51
+
52
+ cfg.PrivateKey = privateKey
53
+ cfg.PublicKey = publicKey
54
+ cfg.ListenPort = utils.IntOrDefault(cfg.ListenPort, DefaultListenPort)
55
+ if len(cfg.OverlayCIDRs) > 0 {
56
+ cfg.OverlayCIDRs, err = utils.NormalizeOverlayCIDRs(cfg.OverlayCIDRs)
57
+ if err != nil {
58
+ return Config{}, fmt.Errorf("normalize overlay cidrs: %w", err)
59
+ }
60
+ }
61
+ if strings.TrimSpace(cfg.Endpoint) == "" {
62
+ cfg.Endpoint = net.JoinHostPort(rootHost, fmt.Sprintf("%d", cfg.ListenPort))
63
+ }
64
+ if strings.TrimSpace(cfg.OverlayIPv4) == "" {
65
+ cfg.OverlayIPv4, err = utils.DeriveWireGuardOverlayIPv4(cfg.PublicKey)
66
+ if err != nil {
67
+ return Config{}, fmt.Errorf("derive overlay ipv4: %w", err)
68
+ }
69
+ }
70
+ if err := utils.ValidateWireGuardEndpoint(cfg.Endpoint); err != nil {
71
+ return Config{}, err
72
+ }
73
+ if err := utils.ValidateOverlayIPv4(cfg.OverlayIPv4); err != nil {
74
+ return Config{}, err
75
+ }
76
+ return cfg, nil
77
+}
78
+
79
+type Overlay struct {
80
+ stack *stack
81
+ listener net.Listener
82
+ server *http.Server
83
+}
84
+
85
+func NewOverlay(cfg Config, handler http.Handler) (*Overlay, error) {
86
+ stack, err := newStack(cfg)
87
+ if err != nil {
88
+ return nil, err
89
+ }
90
+
91
+ listener, err := stack.ListenTCP(DefaultPeerAPIHTTPPort)
92
+ if err != nil {
93
+ _ = stack.Close()
94
+ return nil, err
95
+ }
96
+
97
+ server := &http.Server{
98
+ Handler: handler,
99
+ ReadHeaderTimeout: 10 * time.Second,
100
+ }
101
+
102
+ return &Overlay{
103
+ stack: stack,
104
+ listener: listener,
105
+ server: server,
106
+ }, nil
107
+}
108
+
109
+func (o *Overlay) Serve() error {
110
+ if o == nil || o.server == nil || o.listener == nil {
111
+ return nil
112
+ }
113
+
114
+ err := o.server.Serve(o.listener)
115
+ if errors.Is(err, http.ErrServerClosed) || errors.Is(err, net.ErrClosed) {
116
+ return nil
117
+ }
118
+ return err
119
+}
120
+
121
+func (o *Overlay) Shutdown(ctx context.Context) error {
122
+ if o == nil {
123
+ return nil
124
+ }
125
+
126
+ var shutdownErr error
127
+ if o.server != nil {
128
+ err := o.server.Shutdown(ctx)
129
+ if err != nil && !errors.Is(err, http.ErrServerClosed) {
130
+ shutdownErr = errors.Join(shutdownErr, err)
131
+ }
132
+ }
133
+ if o.listener != nil {
134
+ err := o.listener.Close()
135
+ if err != nil && !errors.Is(err, net.ErrClosed) {
136
+ shutdownErr = errors.Join(shutdownErr, err)
137
+ }
138
+ }
139
+ if o.stack != nil {
140
+ shutdownErr = errors.Join(shutdownErr, o.stack.Close())
141
+ }
142
+ return shutdownErr
143
+}
144
+
145
+func (o *Overlay) Client() *http.Client {
146
+ if o == nil || o.stack == nil {
147
+ return nil
148
+ }
149
+ return &http.Client{
150
+ Transport: &http.Transport{
151
+ DialContext: o.stack.DialContext,
152
+ ForceAttemptHTTP2: false,
153
+ },
154
+ }
155
+}
156
+
157
+func (o *Overlay) Sync(selfRelayID string, snapshot map[string]types.PeerState) 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 {
165
+ peers := make([]types.DesiredPeer, 0, len(snapshot))
166
+ for _, state := range snapshot {
167
+ if state.State != types.PeerStateVerified && state.State != types.PeerStateAdvertised {
168
+ continue
169
+ }
170
+ desc := state.Descriptor
171
+ if desc.RelayID == selfRelayID || !desc.SupportsOverlayPeer {
172
+ continue
173
+ }
174
+ if strings.TrimSpace(desc.WireGuardPublicKey) == "" || strings.TrimSpace(desc.WireGuardEndpoint) == "" || strings.TrimSpace(desc.OverlayIPv4) == "" {
175
+ continue
176
+ }
177
+
178
+ allowedIPs := []string{desc.OverlayIPv4 + "/32"}
179
+ allowedIPs = append(allowedIPs, desc.OverlayCIDRs...)
180
+ peers = append(peers, types.DesiredPeer{
181
+ RelayID: desc.RelayID,
182
+ WireGuardPublicKey: desc.WireGuardPublicKey,
183
+ WireGuardEndpoint: desc.WireGuardEndpoint,
184
+ AllowedIPs: allowedIPs,
185
+ })
186
+ }
187
+ sort.Slice(peers, func(i, j int) bool {
188
+ return peers[i].RelayID < peers[j].RelayID
189
+ })
190
+ return peers
191
+}
portal/wireguard/stack.go
renamed
+29
-104
@@ -5,13 +5,10 @@ import (
5
"errors"
6
"fmt"
7
"net"
8
- "net/http"
8
"net/netip"
10
- "net/url"
9
"strconv"
10
"strings"
11
"sync"
14
- "time"
12
13
"golang.zx2c4.com/wireguard/conn"
14
"golang.zx2c4.com/wireguard/device"
@@ -22,21 +19,13 @@ import (
19
)
20
21
const (
25
- DefaultMTU = 1420
26
- DefaultListenPort = 51820
27
- DefaultPeerAPIHTTPPort = 7777
28
- DefaultPersistentKeepalive = 25
29
- defaultDiscoverRequestTimeout = 15 * time.Second
22
+ DefaultMTU = 1420
23
+ DefaultListenPort = 51820
24
+ DefaultPeerAPIHTTPPort = 7777
25
+ DefaultPersistentKeepalive = 25
26
)
27
32
-type RuntimeConfig struct {
33
- PrivateKey string
34
- Endpoint string
35
- OverlayIPv4 string
36
- MTU int
37
-}
38
-
39
-type Runtime struct {
28
+type stack struct {
29
device *device.Device
30
net *netstack.Net
31
overlayIP netip.Addr
@@ -45,7 +34,7 @@ type Runtime struct {
34
closed bool
35
}
36
48
-func NewRuntime(cfg RuntimeConfig) (*Runtime, error) {
37
+func newStack(cfg Config) (*stack, error) {
38
canonicalPrivateKey, err := utils.NormalizeWireGuardPrivateKey(cfg.PrivateKey)
39
if err != nil {
40
return nil, fmt.Errorf("normalize wireguard private key: %w", err)
@@ -61,12 +50,7 @@ func NewRuntime(cfg RuntimeConfig) (*Runtime, error) {
50
return nil, errors.New("overlay ipv4 must be a valid IPv4 address")
51
}
52
64
- mtu := cfg.MTU
65
- if mtu <= 0 {
66
- mtu = DefaultMTU
67
- }
68
-
69
- tunDevice, network, err := netstack.CreateNetTUN([]netip.Addr{overlayIP}, nil, mtu)
53
+ tunDevice, network, err := netstack.CreateNetTUN([]netip.Addr{overlayIP}, nil, DefaultMTU)
54
if err != nil {
55
return nil, fmt.Errorf("create netstack tun: %w", err)
56
}
@@ -91,26 +75,26 @@ func NewRuntime(cfg RuntimeConfig) (*Runtime, error) {
75
return nil, fmt.Errorf("bring wireguard device up: %w", err)
76
}
77
94
- return &Runtime{
78
+ return &stack{
79
device: wgDevice,
80
net: network,
81
overlayIP: overlayIP,
82
}, nil
83
}
84
101
-func (r *Runtime) ListenTCP(port int) (net.Listener, error) {
102
- if r == nil || r.net == nil {
103
- return nil, errors.New("wireguard runtime is not initialized")
85
+func (s *stack) ListenTCP(port int) (net.Listener, error) {
86
+ if s == nil || s.net == nil {
87
+ return nil, errors.New("wireguard is not initialized")
88
}
105
- return r.net.ListenTCP(&net.TCPAddr{
106
- IP: net.ParseIP(r.overlayIP.String()),
89
+ return s.net.ListenTCP(&net.TCPAddr{
90
+ IP: net.ParseIP(s.overlayIP.String()),
91
Port: port,
92
})
93
}
94
111
-func (r *Runtime) DialContext(ctx context.Context, network, address string) (net.Conn, error) {
112
- if r == nil || r.net == nil {
113
- return nil, errors.New("wireguard runtime is not initialized")
95
+func (s *stack) DialContext(ctx context.Context, network, address string) (net.Conn, error) {
96
+ if s == nil || s.net == nil {
97
+ return nil, errors.New("wireguard is not initialized")
98
}
99
switch network {
100
case "tcp", "tcp4", "tcp6":
@@ -130,71 +114,12 @@ func (r *Runtime) DialContext(ctx context.Context, network, address string) (net
114
if err != nil || port <= 0 || port > 65535 {
115
return nil, errors.New("invalid tcp port")
116
}
133
- return r.net.DialContextTCPAddrPort(ctx, netip.AddrPortFrom(ip, uint16(port)))
134
-}
135
-
136
-func (r *Runtime) Discover(ctx context.Context, overlayIPv4 string, port int, req types.DiscoverRequest) (types.DiscoverResponse, error) {
137
- if r == nil {
138
- return types.DiscoverResponse{}, errors.New("wireguard runtime is not initialized")
139
- }
140
- if port == 0 {
141
- port = DefaultPeerAPIHTTPPort
142
- }
143
- ip, err := netip.ParseAddr(strings.TrimSpace(overlayIPv4))
144
- if err != nil || !ip.Is4() {
145
- return types.DiscoverResponse{}, errors.New("overlay ipv4 must be a valid IPv4 address")
146
- }
147
-
148
- baseURL := &url.URL{
149
- Scheme: "http",
150
- Host: net.JoinHostPort(ip.String(), strconv.Itoa(port)),
151
- Path: types.PathDiscovery,
152
- }
153
- query := baseURL.Query()
154
- if req.RootHost != "" {
155
- query.Set("root_host", req.RootHost)
156
- }
157
- if req.Name != "" {
158
- query.Set("name", req.Name)
159
- }
160
- baseURL.RawQuery = query.Encode()
161
-
162
- httpClient := &http.Client{
163
- Transport: &http.Transport{
164
- DialContext: r.DialContext,
165
- ForceAttemptHTTP2: false,
166
- },
167
- Timeout: defaultDiscoverRequestTimeout,
168
- }
169
-
170
- httpReq, err := http.NewRequestWithContext(ctx, http.MethodGet, baseURL.String(), nil)
171
- if err != nil {
172
- return types.DiscoverResponse{}, err
173
- }
174
-
175
- resp, err := httpClient.Do(httpReq)
176
- if err != nil {
177
- return types.DiscoverResponse{}, err
178
- }
179
- defer resp.Body.Close()
180
-
181
- if resp.StatusCode < http.StatusOK || resp.StatusCode >= http.StatusMultipleChoices {
182
- return types.DiscoverResponse{}, utils.DecodeAPIRequestError(resp)
183
- }
184
-
185
- envelope, err := utils.DecodeAPIEnvelope[types.DiscoverResponse](resp.Body)
186
- if err != nil {
187
- return types.DiscoverResponse{}, fmt.Errorf("decode response: %w", err)
188
- }
189
- if !envelope.OK {
190
- return types.DiscoverResponse{}, utils.NewAPIRequestError(resp.StatusCode, envelope.Error)
191
- }
192
- return envelope.Data, nil
117
+ return s.net.DialContextTCPAddrPort(ctx, netip.AddrPortFrom(ip, uint16(port)))
118
}
119
195
-func (r *Runtime) ApplyPeers(peers []types.DesiredPeer) error {
196
- if r == nil || r.device == nil {
197
- return errors.New("wireguard runtime is not initialized")
120
+func (s *stack) ApplyPeers(peers []types.DesiredPeer) error {
121
+ if s == nil || s.device == nil {
122
+ return errors.New("wireguard is not initialized")
123
}
124
125
var builder strings.Builder
@@ -227,22 +152,22 @@ func (r *Runtime) ApplyPeers(peers []types.DesiredPeer) error {
152
}
153
}
154
230
- return r.device.IpcSet(builder.String())
155
+ return s.device.IpcSet(builder.String())
156
}
157
233
-func (r *Runtime) Close() error {
234
- if r == nil || r.device == nil {
158
+func (s *stack) Close() error {
159
+ if s == nil || s.device == nil {
160
return nil
161
}
162
238
- r.mu.Lock()
239
- if r.closed {
240
- r.mu.Unlock()
163
+ s.mu.Lock()
164
+ if s.closed {
165
+ s.mu.Unlock()
166
return nil
167
}
243
- r.closed = true
244
- device := r.device
245
- r.mu.Unlock()
168
+ s.closed = true
169
+ device := s.device
170
+ s.mu.Unlock()
171
172
device.Close()
173
<-device.Wait()
portal/wireguard/stack_test.go
renamed
+4
-4
@@ -32,7 +32,7 @@ func TestNormalizePrivateKeyAndPublicKeyFromPrivate(t *testing.T) {
32
}
33
}
34
35
-func TestRuntimeStartAndClose(t *testing.T) {
35
+func TestStackStartAndClose(t *testing.T) {
36
t.Parallel()
37
38
privateKey, err := utils.NormalizeWireGuardPrivateKey("2222222222222222222222222222222222222222222222222222222222222222")
@@ -41,16 +41,16 @@ func TestRuntimeStartAndClose(t *testing.T) {
41
}
42
43
port := reserveUDPPort(t)
44
- runtime, err := NewRuntime(RuntimeConfig{
44
+ stack, err := newStack(Config{
45
PrivateKey: privateKey,
46
Endpoint: net.JoinHostPort("127.0.0.1", port),
47
OverlayIPv4: "10.77.0.1",
48
})
49
if err != nil {
50
- t.Fatalf("NewRuntime() error = %v", err)
50
+ t.Fatalf("newStack() error = %v", err)
51
}
52
t.Cleanup(func() {
53
- if err := runtime.Close(); err != nil {
53
+ if err := stack.Close(); err != nil {
54
t.Fatalf("Close() error = %v", err)
55
}
56
})
sdk/expose.go
+168
-190
@@ -70,7 +70,7 @@ func Expose(ctx context.Context, cfg ExposeConfig) (*Exposure, error) {
70
return nil, err
71
}
72
73
- identity, err := discovery.ResolveIdentity(cfg.OwnerPrivateKey)
73
+ identity, err := utils.ResolveSecp256k1Identity(cfg.OwnerPrivateKey)
74
if err != nil {
75
return nil, fmt.Errorf("resolve owner identity: %w", err)
76
}
@@ -106,7 +106,7 @@ func Expose(ctx context.Context, cfg ExposeConfig) (*Exposure, error) {
106
}
107
108
if len(relayURLs) > 0 {
109
- if _, err := exposure.applyRelayURLs(relayURLs, true); err != nil {
109
+ if _, err := exposure.setRelayURLs(relayURLs, true); err != nil {
110
_ = exposure.Close()
111
return nil, err
112
}
@@ -124,10 +124,13 @@ func Expose(ctx context.Context, cfg ExposeConfig) (*Exposure, error) {
124
}()
125
126
if len(relayURLs) > 0 {
127
+ exposure.mu.RLock()
128
+ activeRelayURLs := append([]string(nil), exposure.activeRelayURLs...)
129
+ exposure.mu.RUnlock()
130
log.Info().
131
Str("release_version", types.ReleaseVersion).
129
- Int("relay_count", len(exposure.ActiveRelayURLs())).
130
- Strs("relays", exposure.ActiveRelayURLs()).
132
+ Int("relay_count", len(activeRelayURLs)).
133
+ Strs("relays", activeRelayURLs).
134
Msg("exposure relay started")
135
}
136
@@ -196,37 +199,12 @@ func (e *Exposure) Addr() net.Addr {
199
return listenerAddr("portal:exposure")
200
}
201
199
-func (e *Exposure) PublicURLs() []string {
200
- listeners := e.listenersOrdered()
201
- if len(listeners) == 0 {
202
- return nil
203
- }
204
-
205
- out := make([]string, 0, len(listeners))
206
- seen := make(map[string]struct{})
207
- for _, listener := range listeners {
208
- if listener == nil {
209
- continue
210
- }
211
- rawURL := listener.PublicURL()
212
- if rawURL == "" {
213
- continue
214
- }
215
- if _, ok := seen[rawURL]; ok {
216
- continue
217
- }
218
- seen[rawURL] = struct{}{}
219
- out = append(out, rawURL)
220
- }
221
- if len(out) == 0 {
222
- return nil
223
- }
224
- return out
225
-}
226
-
202
func (e *Exposure) RunHTTP(ctx context.Context, handler http.Handler, localAddr string) error {
203
var relayListener net.Listener
229
- if len(e.ActiveRelayURLs()) > 0 {
204
+ e.mu.RLock()
205
+ hasActiveRelays := len(e.activeRelayURLs) > 0
206
+ e.mu.RUnlock()
207
+ if hasActiveRelays {
208
relayListener = e
209
}
210
return RunHTTP(ctx, relayListener, handler, localAddr)
@@ -340,7 +318,16 @@ func (e *Exposure) Close() error {
318
e.cancel()
319
}
320
343
- listeners := e.listenersOrdered()
321
+ e.mu.RLock()
322
+ listeners := make([]*Listener, 0, len(e.listeners))
323
+ activeRelayURLs := append([]string(nil), e.activeRelayURLs...)
324
+ for _, relayURL := range e.activeRelayURLs {
325
+ if listener, ok := e.listeners[relayURL]; ok {
326
+ listeners = append(listeners, listener)
327
+ }
328
+ }
329
+ e.mu.RUnlock()
330
+
331
for _, listener := range listeners {
332
if listener != nil {
333
closeErr = errors.Join(closeErr, listener.Close())
@@ -349,49 +336,54 @@ func (e *Exposure) Close() error {
336
337
event := log.Info().
338
Int("relay_count", len(listeners)).
352
- Strs("relays", e.ActiveRelayURLs())
339
+ Strs("relays", activeRelayURLs)
340
if closeErr != nil {
341
event = log.Warn().
342
Err(closeErr).
343
Int("relay_count", len(listeners)).
357
- Strs("relays", e.ActiveRelayURLs())
344
+ Strs("relays", activeRelayURLs)
345
}
346
event.Msg("exposure closed")
347
})
348
return closeErr
349
}
350
364
-func (e *Exposure) applyRelayURLs(relayURLs []string, failOnError bool) ([]string, error) {
351
+func (e *Exposure) setRelayURLs(relayURLs []string, failOnError bool) ([]string, error) {
352
if len(relayURLs) == 0 {
353
return nil, nil
354
}
355
369
- snapshot := append([]string(nil), relayURLs...)
370
-
356
+ relayURLs = utils.FilterRelayURLs(append([]string(nil), relayURLs...), e.bannedRelayURLs)
357
e.mu.Lock()
372
- snapshot = utils.FilterRelayURLs(snapshot, e.bannedRelayURLs)
358
existing := make(map[string]struct{}, len(e.knownRelayURLs))
359
for _, relayURL := range e.knownRelayURLs {
360
existing[relayURL] = struct{}{}
361
}
377
- if strings.Join(e.knownRelayURLs, "\x00") != strings.Join(snapshot, "\x00") {
378
- e.knownRelayURLs = snapshot
379
- }
380
- if strings.Join(e.activeRelayURLs, "\x00") != strings.Join(snapshot, "\x00") {
381
- e.activeRelayURLs = snapshot
362
+
363
+ added := make([]string, 0, len(relayURLs))
364
+ missing := make([]string, 0)
365
+ for _, relayURL := range relayURLs {
366
+ if _, ok := existing[relayURL]; !ok {
367
+ added = append(added, relayURL)
368
+ }
369
+ if _, ok := e.listeners[relayURL]; !ok {
370
+ missing = append(missing, relayURL)
371
+ }
372
}
373
+ e.knownRelayURLs = append([]string(nil), relayURLs...)
374
+ e.activeRelayURLs = append([]string(nil), relayURLs...)
375
e.mu.Unlock()
376
385
- added := make([]string, 0, len(snapshot))
386
- for _, relayURL := range snapshot {
387
- if _, ok := existing[relayURL]; ok {
377
+ for _, relayURL := range missing {
378
+ listener, err := e.newListener(relayURL)
379
+ if err != nil {
380
+ if failOnError {
381
+ return nil, fmt.Errorf("listen %q: %w", relayURL, err)
382
+ }
383
+ log.Warn().Err(err).Str("relay_url", relayURL).Msg("add relay listener")
384
continue
385
}
390
- added = append(added, relayURL)
391
- }
392
-
393
- if err := e.syncListeners(failOnError); err != nil {
394
- return nil, err
386
+ e.installListener(relayURL, listener)
387
}
388
return added, nil
389
}
@@ -411,35 +403,14 @@ func (e *Exposure) banRelayURL(relayURL string) {
403
Msg("relay banned by mitm detection")
404
}
405
414
-func (e *Exposure) syncListeners(failOnError bool) error {
415
- e.mu.Lock()
416
- missing := make([]string, 0)
417
- for _, relayURL := range e.activeRelayURLs {
418
- if _, ok := e.listeners[relayURL]; ok {
419
- continue
420
- }
421
- missing = append(missing, relayURL)
422
- }
423
- e.mu.Unlock()
424
-
425
- for _, relayURL := range missing {
426
- listener, err := e.newListener(relayURL)
427
- if err != nil {
428
- if failOnError {
429
- return fmt.Errorf("listen %q: %w", relayURL, err)
430
- }
431
- log.Warn().Err(err).Str("relay_url", relayURL).Msg("add relay listener")
432
- continue
433
- }
434
- e.installListener(relayURL, listener)
435
- }
436
- return nil
437
-}
438
-
406
func (e *Exposure) newListener(relayURL string) (*Listener, error) {
407
bootstraps := []string(nil)
408
if e.discoveryEnabled {
442
- bootstraps = e.KnownRelayURLs()
409
+ e.mu.RLock()
410
+ if len(e.knownRelayURLs) > 0 {
411
+ bootstraps = append([]string(nil), e.knownRelayURLs...)
412
+ }
413
+ e.mu.RUnlock()
414
}
415
416
cfg := ListenerConfig{
@@ -462,12 +433,15 @@ func (e *Exposure) installListener(relayURL string, listener *Listener) {
433
434
shouldClose := false
435
e.mu.Lock()
465
- if e.closed() {
466
- shouldClose = true
467
- } else if _, exists := e.listeners[relayURL]; exists {
436
+ select {
437
+ case <-e.done:
438
shouldClose = true
469
- } else {
470
- e.listeners[relayURL] = listener
439
+ default:
440
+ if _, exists := e.listeners[relayURL]; exists {
441
+ shouldClose = true
442
+ } else {
443
+ e.listeners[relayURL] = listener
444
+ }
445
}
446
e.mu.Unlock()
447
@@ -481,45 +455,15 @@ func (e *Exposure) installListener(relayURL string, listener *Listener) {
455
if e.udpEnabled {
456
go func() {
457
relayURL := listener.api.baseURL.String()
484
- if err := listener.WaitRegistered(context.Background()); err != nil {
485
- switch {
486
- case e.closed():
487
- return
488
- case errors.Is(err, net.ErrClosed), errors.Is(err, context.Canceled):
489
- return
490
- default:
491
- log.Warn().
492
- Err(err).
493
- Str("relay_url", relayURL).
494
- Msg("attach datagram plane failed")
495
- return
496
- }
497
- }
498
- if listener.UDPAddr() == "" {
499
- if !e.closed() && !listener.closed() {
500
- log.Warn().
501
- Str("relay_url", relayURL).
502
- Msg("attach datagram plane failed")
503
- }
504
- return
505
- }
506
-
507
- ticker := time.NewTicker(50 * time.Millisecond)
508
- defer ticker.Stop()
509
- for listener.datagram == nil || !listener.datagram.Connected() {
510
- select {
511
- case <-e.done:
512
- return
513
- case <-listener.doneCh:
514
- return
515
- case <-ticker.C:
516
- }
517
- }
518
-
458
for {
520
- frame, err := listener.datagram.Accept(listener.doneCh)
459
+ frame, err := listener.AcceptDatagram()
460
if err != nil {
522
- if e.closed() || errors.Is(err, net.ErrClosed) {
461
+ select {
462
+ case <-e.done:
463
+ return
464
+ default:
465
+ }
466
+ if errors.Is(err, net.ErrClosed) {
467
return
468
}
469
log.Warn().
@@ -530,11 +474,6 @@ func (e *Exposure) installListener(relayURL string, listener *Listener) {
474
return
475
}
476
533
- frame.Payload = append([]byte(nil), frame.Payload...)
534
- frame.LeaseID = listener.LeaseID()
535
- frame.RelayURL = relayURL
536
- frame.UDPAddr = listener.UDPAddr()
537
-
477
select {
478
case <-e.done:
479
return
@@ -545,19 +484,6 @@ func (e *Exposure) installListener(relayURL string, listener *Listener) {
484
}
485
}
486
548
-func (e *Exposure) listenersOrdered() []*Listener {
549
- e.mu.RLock()
550
- defer e.mu.RUnlock()
551
-
552
- out := make([]*Listener, 0, len(e.listeners))
553
- for _, relayURL := range e.activeRelayURLs {
554
- if listener, ok := e.listeners[relayURL]; ok {
555
- out = append(out, listener)
556
- }
557
- }
558
- return out
559
-}
560
-
487
func (e *Exposure) runListenerAcceptLoop(listener *Listener) {
488
if listener == nil {
489
return
@@ -646,21 +572,13 @@ func (e *Exposure) SendDatagram(frame types.DatagramFrame) error {
572
return net.ErrClosed
573
}
574
649
- relayURL := strings.TrimSpace(frame.RelayURL)
650
- if relayURL == "" {
651
- return errors.New("relay url is required")
652
- }
653
-
575
e.mu.RLock()
655
- listener := e.listeners[relayURL]
576
+ listener := e.listeners[frame.RelayURL]
577
e.mu.RUnlock()
657
- if listener == nil || listener.datagram == nil {
578
+ if listener == nil {
579
return net.ErrClosed
580
}
660
- if leaseID := strings.TrimSpace(frame.LeaseID); leaseID != "" && leaseID != listener.LeaseID() {
661
- return errors.New("datagram frame targets stale lease")
662
- }
663
- return listener.datagram.Send(frame.FlowID, frame.Payload)
581
+ return listener.SendDatagram(frame)
582
}
583
584
func (e *Exposure) WaitDatagramReady(ctx context.Context) ([]string, error) {
@@ -672,7 +590,15 @@ func (e *Exposure) WaitDatagramReady(ctx context.Context) ([]string, error) {
590
defer ticker.Stop()
591
592
for {
675
- listeners := e.listenersOrdered()
593
+ e.mu.RLock()
594
+ listeners := make([]*Listener, 0, len(e.listeners))
595
+ for _, relayURL := range e.activeRelayURLs {
596
+ if listener, ok := e.listeners[relayURL]; ok {
597
+ listeners = append(listeners, listener)
598
+ }
599
+ }
600
+ e.mu.RUnlock()
601
+
602
addrs := make([]string, 0, len(listeners))
603
seen := make(map[string]struct{})
604
resolvedWithoutDatagram := true
@@ -681,23 +607,15 @@ func (e *Exposure) WaitDatagramReady(ctx context.Context) ([]string, error) {
607
continue
608
}
609
684
- udpAddr := listener.UDPAddr()
685
- if listener.datagram != nil && listener.datagram.Connected() && udpAddr != "" {
610
+ udpAddr, ready, pending := listener.DatagramReady()
611
+ if ready {
612
if _, ok := seen[udpAddr]; !ok {
613
seen[udpAddr] = struct{}{}
614
addrs = append(addrs, udpAddr)
615
}
616
}
691
-
692
- select {
693
- case <-listener.registered:
694
- if udpAddr != "" {
695
- resolvedWithoutDatagram = false
696
- }
697
- default:
698
- if !listener.closed() {
699
- resolvedWithoutDatagram = false
700
- }
617
+ if pending {
618
+ resolvedWithoutDatagram = false
619
}
620
}
621
if len(addrs) > 0 {
@@ -717,15 +635,6 @@ func (e *Exposure) WaitDatagramReady(ctx context.Context) ([]string, error) {
635
}
636
}
637
720
-func (e *Exposure) closed() bool {
721
- select {
722
- case <-e.done:
723
- return true
724
- default:
725
- return false
726
- }
727
-}
728
-
638
func (e *Exposure) monitorStartupCounts() {
639
ticker := time.NewTicker(time.Second)
640
defer ticker.Stop()
@@ -733,11 +642,21 @@ func (e *Exposure) monitorStartupCounts() {
642
firstRun := true
643
644
for {
645
+ e.mu.RLock()
646
+ listeners := make([]*Listener, 0, len(e.listeners))
647
+ for _, relayURL := range e.activeRelayURLs {
648
+ if listener, ok := e.listeners[relayURL]; ok {
649
+ listeners = append(listeners, listener)
650
+ }
651
+ }
652
+ bannedCount := len(e.bannedRelayURLs)
653
+ e.mu.RUnlock()
654
+
655
readyCount, inactiveCount := 0, 0
656
activated := make([]string, 0)
657
deactivated := make([]string, 0)
658
740
- for _, listener := range e.listenersOrdered() {
659
+ for _, listener := range listeners {
660
if listener == nil {
661
continue
662
}
@@ -761,7 +680,6 @@ func (e *Exposure) monitorStartupCounts() {
680
}
681
682
if firstRun || len(activated) > 0 || len(deactivated) > 0 {
764
- bannedCount := len(e.BannedRelayURLs())
683
event := log.Info().
684
Int("banned", bannedCount).
685
Int("inactive", inactiveCount).
@@ -792,30 +710,80 @@ func (e *Exposure) runDiscoveryLoop(ctx context.Context) {
710
discoveryFailed := false
711
712
for {
795
- peers := e.KnownRelayURLs()
796
- if len(peers) > 0 {
797
- relayURLs, err := discovery.DiscoverBootstraps(ctx, peers, types.DiscoverRequest{}, e.rootCAPEM)
713
+ e.mu.RLock()
714
+ knownRelayURLs := append([]string(nil), e.knownRelayURLs...)
715
+ e.mu.RUnlock()
716
+
717
+ peers, err := utils.ExcludeLocalRelayURLs(knownRelayURLs...)
718
+ if err == nil && len(peers) > 0 {
719
+ relayURLs := append([]string(nil), peers...)
720
+ var discoverErr error
721
+
722
+ for _, peer := range peers {
723
+ resp, err := discovery.Discover(ctx, peer, types.DiscoverRequest{}, e.rootCAPEM, nil)
724
+ if err != nil {
725
+ discoverErr = errors.Join(discoverErr, fmt.Errorf("discover %q: %w", peer, err))
726
+ continue
727
+ }
728
+
729
+ now := time.Now().UTC()
730
+ self, descriptors, err := discovery.ValidateResponse(resp, now)
731
+ if err != nil {
732
+ if self.RelayID == "" {
733
+ discoverErr = errors.Join(discoverErr, fmt.Errorf("validate %q self descriptor: %w", peer, err))
734
+ continue
735
+ }
736
+ discoverErr = errors.Join(discoverErr, fmt.Errorf("validate %q peer descriptors: %w", peer, err))
737
+ }
738
+
739
+ urls := make([]string, 0, 1+len(descriptors))
740
+ if apiURL := strings.TrimSpace(self.APIHTTPSAddr); apiURL != "" {
741
+ urls = append(urls, apiURL)
742
+ }
743
+ for _, descriptor := range descriptors {
744
+ if apiURL := strings.TrimSpace(descriptor.APIHTTPSAddr); apiURL != "" {
745
+ urls = append(urls, apiURL)
746
+ }
747
+ }
748
+
749
+ discoveredRelayURLs, err := utils.ExcludeLocalRelayURLs(urls...)
750
+ if err != nil {
751
+ discoverErr = errors.Join(discoverErr, fmt.Errorf("extract %q relay urls: %w", peer, err))
752
+ continue
753
+ }
754
+ relayURLs, err = utils.MergeRelayURLs(relayURLs, nil, discoveredRelayURLs)
755
+ if err != nil {
756
+ discoverErr = errors.Join(discoverErr, fmt.Errorf("merge %q relay urls: %w", peer, err))
757
+ continue
758
+ }
759
+ }
760
+
761
+ err = discoverErr
762
switch {
763
case err == nil:
800
- if discoveryFailed {
801
- log.Info().
802
- Int("peer_count", len(peers)).
803
- Msg("relay discovery recovered")
804
- }
764
+ recovered := discoveryFailed
765
discoveryFailed = false
806
- added, err := e.applyRelayURLs(relayURLs, false)
766
+ added, err := e.setRelayURLs(relayURLs, false)
767
if err != nil {
768
log.Warn().
769
Err(err).
770
Int("relay_count", len(peers)).
811
- Msg("apply discovered relay urls failed")
812
- } else if len(added) > 0 {
813
- log.Info().
771
+ Msg("discover relay urls failed")
772
+ } else if recovered || len(added) > 0 {
773
+ e.mu.RLock()
774
+ totalKnownRelayCount := len(e.knownRelayURLs)
775
+ e.mu.RUnlock()
776
+ event := log.Info().
777
Int("peer_count", len(peers)).
815
- Int("added_count", len(added)).
816
- Int("total_known_relay_count", len(e.KnownRelayURLs())).
817
- Strs("added_relays", added).
818
- Msg("discovery relays updated")
778
+ Int("total_known_relay_count", totalKnownRelayCount)
779
+ if recovered {
780
+ event = event.Bool("recovered", true)
781
+ }
782
+ if len(added) > 0 {
783
+ event = event.Int("added_count", len(added)).
784
+ Strs("added_relays", added)
785
+ }
786
+ event.Msg("discovery relays updated")
787
}
788
case ctx.Err() != nil:
789
return
@@ -828,6 +796,16 @@ func (e *Exposure) runDiscoveryLoop(ctx context.Context) {
796
}
797
discoveryFailed = true
798
}
799
+ } else if err != nil {
800
+ if ctx.Err() != nil {
801
+ return
802
+ }
803
+ if !discoveryFailed {
804
+ log.Debug().
805
+ Err(err).
806
+ Msg("discover relay urls failed")
807
+ }
808
+ discoveryFailed = true
809
}
810
811
select {
sdk/expose_test.go
+3
-3
@@ -51,7 +51,7 @@ func TestExposureBanRelayURLMovesRelay(t *testing.T) {
51
}
52
}
53
54
-func TestExposureApplyRelayURLsSkipsBannedRelay(t *testing.T) {
54
+func TestExposureSetRelayURLsSkipsBannedRelay(t *testing.T) {
55
const (
56
relayA = "https://relay-a.example"
57
relayB = "https://relay-b.example"
@@ -64,9 +64,9 @@ func TestExposureApplyRelayURLsSkipsBannedRelay(t *testing.T) {
64
},
65
}
66
67
- added, err := exposure.applyRelayURLs([]string{relayA, relayB}, false)
67
+ added, err := exposure.setRelayURLs([]string{relayA, relayB}, false)
68
if err != nil {
69
- t.Fatalf("applyRelayURLs() error = %v", err)
69
+ t.Fatalf("setRelayURLs() error = %v", err)
70
}
71
if len(added) != 1 || added[0] != relayA {
72
t.Fatalf("added relay urls = %v, want [%q]", added, relayA)
sdk/listener.go
+59
-18
@@ -265,6 +265,65 @@ func (l *Listener) Accept() (net.Conn, error) {
265
}
266
}
267
268
+func (l *Listener) AcceptDatagram() (types.DatagramFrame, error) {
269
+ if l == nil || l.datagram == nil {
270
+ return types.DatagramFrame{}, net.ErrClosed
271
+ }
272
+
273
+ frame, err := l.datagram.Accept(l.doneCh)
274
+ if err != nil {
275
+ return types.DatagramFrame{}, err
276
+ }
277
+
278
+ frame.Payload = append([]byte(nil), frame.Payload...)
279
+ l.mu.Lock()
280
+ frame.LeaseID = l.leaseID
281
+ frame.UDPAddr = l.udpAddr
282
+ if l.api != nil && l.api.baseURL != nil {
283
+ frame.RelayURL = l.api.baseURL.String()
284
+ }
285
+ l.mu.Unlock()
286
+ return frame, nil
287
+}
288
+
289
+func (l *Listener) SendDatagram(frame types.DatagramFrame) error {
290
+ if l == nil || l.datagram == nil {
291
+ return net.ErrClosed
292
+ }
293
+
294
+ l.mu.Lock()
295
+ leaseID := l.leaseID
296
+ datagram := l.datagram
297
+ l.mu.Unlock()
298
+
299
+ if leaseID == "" || datagram == nil {
300
+ return net.ErrClosed
301
+ }
302
+ if frameLeaseID := strings.TrimSpace(frame.LeaseID); frameLeaseID != "" && frameLeaseID != leaseID {
303
+ return errors.New("datagram frame targets stale lease")
304
+ }
305
+ return datagram.Send(frame.FlowID, frame.Payload)
306
+}
307
+
308
+func (l *Listener) DatagramReady() (string, bool, bool) {
309
+ if l == nil || l.datagram == nil {
310
+ return "", false, false
311
+ }
312
+
313
+ l.mu.Lock()
314
+ udpAddr := l.udpAddr
315
+ datagram := l.datagram
316
+ l.mu.Unlock()
317
+
318
+ ready := datagram != nil && datagram.Connected() && udpAddr != ""
319
+ select {
320
+ case <-l.registered:
321
+ return udpAddr, ready, udpAddr != "" && !ready
322
+ default:
323
+ return udpAddr, ready, !l.closed()
324
+ }
325
+}
326
+
327
func (l *Listener) Addr() net.Addr {
328
l.mu.Lock()
329
defer l.mu.Unlock()
@@ -320,12 +379,6 @@ func (l *Listener) PublicURL() string {
379
}).String()
380
}
381
323
-func (l *Listener) UDPAddr() string {
324
- l.mu.Lock()
325
- defer l.mu.Unlock()
326
- return l.udpAddr
327
-}
328
-
382
func (l *Listener) currentDatagramState() (transport.ClientDatagramState, bool) {
383
if l.datagram == nil {
384
return transport.ClientDatagramState{}, false
@@ -462,18 +515,6 @@ func (l *Listener) registerAndConfigure(ctx context.Context, registerBootstraps
515
return nil
516
}
517
465
-// WaitRegistered blocks until the first successful lease registration or context cancellation.
466
-func (l *Listener) WaitRegistered(ctx context.Context) error {
467
- select {
468
- case <-l.registered:
469
- return nil
470
- case <-l.doneCh:
471
- return net.ErrClosed
472
- case <-ctx.Done():
473
- return ctx.Err()
474
- }
475
-}
476
-
518
func (l *Listener) retryOrClose(ctx context.Context, operation string, err error, retries int) bool {
519
if ctx.Err() != nil {
520
return false
sdk/sdk_test.go
+4
-4
@@ -9,8 +9,8 @@ import (
9
"testing"
10
"time"
11
12
- "github.com/gosuda/portal/v2/portal/discovery"
12
"github.com/gosuda/portal/v2/types"
13
+ "github.com/gosuda/portal/v2/utils"
14
)
15
16
func TestNewListenerRegistersLeaseWithMainContract(t *testing.T) {
@@ -121,9 +121,9 @@ func TestExposeNoRelayInputs(t *testing.T) {
121
122
func TestExposeResolvesOwnerPrivateKey(t *testing.T) {
123
ownerPrivateKey := strings.Repeat("11", 32)
124
- identity, err := discovery.ResolveIdentity(ownerPrivateKey)
124
+ identity, err := utils.ResolveSecp256k1Identity(ownerPrivateKey)
125
if err != nil {
126
- t.Fatalf("ResolveIdentity() error = %v", err)
126
+ t.Fatalf("ResolveSecp256k1Identity() error = %v", err)
127
}
128
129
registerReqCh := make(chan types.RegisterRequest, 1)
@@ -262,7 +262,7 @@ func TestExposeGeneratesOwnerAddressWithoutPrivateKey(t *testing.T) {
262
if registerReq.OwnerAddress == "" {
263
t.Fatal("register request OwnerAddress = empty, want generated address")
264
}
265
- if _, err := discovery.NormalizeEVMAddress(registerReq.OwnerAddress); err != nil {
265
+ if _, err := utils.NormalizeEVMAddress(registerReq.OwnerAddress); err != nil {
266
t.Fatalf("register request OwnerAddress = %q, want valid EVM address: %v", registerReq.OwnerAddress, err)
267
}
268
}
utils/crypto.go
new
+393
@@ -0,0 +1,393 @@
1
+package utils
2
+
3
+import (
4
+ "crypto/sha256"
5
+ "encoding/base64"
6
+ "encoding/hex"
7
+ "errors"
8
+ "fmt"
9
+ "net"
10
+ "net/netip"
11
+ "sort"
12
+ "strconv"
13
+ "strings"
14
+
15
+ "github.com/decred/dcrd/dcrec/secp256k1/v4"
16
+ secp256k1ecdsa "github.com/decred/dcrd/dcrec/secp256k1/v4/ecdsa"
17
+ "golang.org/x/crypto/curve25519"
18
+ "golang.org/x/crypto/sha3"
19
+)
20
+
21
+type Secp256k1Identity struct {
22
+ Generated bool `json:"generated,omitempty"`
23
+ Address string `json:"address"`
24
+ PublicKey string `json:"public_key"`
25
+ PrivateKey string `json:"private_key"`
26
+}
27
+
28
+func AddressFromCompressedPublicKeyHex(rawPublicKey string) (string, error) {
29
+ publicKeyHex := strings.TrimSpace(rawPublicKey)
30
+ if publicKeyHex == "" {
31
+ return "", errors.New("public key is required")
32
+ }
33
+ if strings.HasPrefix(strings.ToLower(publicKeyHex), "0x") {
34
+ publicKeyHex = publicKeyHex[2:]
35
+ }
36
+
37
+ decoded, err := hex.DecodeString(publicKeyHex)
38
+ if err != nil {
39
+ return "", errors.New("public key must be hex encoded")
40
+ }
41
+
42
+ publicKey, err := secp256k1.ParsePubKey(decoded)
43
+ if err != nil {
44
+ return "", errors.New("invalid secp256k1 public key")
45
+ }
46
+
47
+ uncompressed := publicKey.SerializeUncompressed()
48
+ if len(uncompressed) != 65 || uncompressed[0] != 0x04 {
49
+ return "", errors.New("invalid uncompressed secp256k1 public key")
50
+ }
51
+
52
+ hasher := sha3.NewLegacyKeccak256()
53
+ _, _ = hasher.Write(uncompressed[1:])
54
+ hash := hasher.Sum(nil)
55
+
56
+ return NormalizeEVMAddress("0x" + hex.EncodeToString(hash[len(hash)-20:]))
57
+}
58
+
59
+func NormalizeEVMAddress(raw string) (string, error) {
60
+ trimmed := strings.TrimSpace(raw)
61
+ if trimmed == "" {
62
+ return "", errors.New("address is required")
63
+ }
64
+ if !strings.HasPrefix(strings.ToLower(trimmed), "0x") {
65
+ return "", errors.New("address must start with 0x")
66
+ }
67
+
68
+ hexPart := trimmed[2:]
69
+ if len(hexPart) != 40 {
70
+ return "", errors.New("address must be 20 bytes")
71
+ }
72
+ if _, err := hex.DecodeString(hexPart); err != nil {
73
+ return "", errors.New("address must be hex encoded")
74
+ }
75
+
76
+ lowerHex := strings.ToLower(hexPart)
77
+ hasher := sha3.NewLegacyKeccak256()
78
+ _, _ = hasher.Write([]byte(lowerHex))
79
+ hash := hasher.Sum(nil)
80
+
81
+ var builder strings.Builder
82
+ builder.Grow(len(lowerHex))
83
+ for idx, ch := range lowerHex {
84
+ if ch >= '0' && ch <= '9' {
85
+ builder.WriteRune(ch)
86
+ continue
87
+ }
88
+
89
+ nibble := hash[idx/2]
90
+ if idx%2 == 0 {
91
+ nibble >>= 4
92
+ } else {
93
+ nibble &= 0x0f
94
+ }
95
+ if nibble > 7 {
96
+ builder.WriteRune(ch - ('a' - 'A'))
97
+ continue
98
+ }
99
+ builder.WriteRune(ch)
100
+ }
101
+
102
+ checksummed := builder.String()
103
+ if hexPart != lowerHex && hexPart != strings.ToUpper(hexPart) && hexPart != checksummed {
104
+ return "", errors.New("address checksum is invalid")
105
+ }
106
+ return "0x" + checksummed, nil
107
+}
108
+
109
+func ResolveSecp256k1Identity(rawPrivateKey string) (Secp256k1Identity, error) {
110
+ privateKeyHex := strings.TrimSpace(rawPrivateKey)
111
+ generated := false
112
+ if privateKeyHex == "" {
113
+ privateKey, err := secp256k1.GeneratePrivateKey()
114
+ if err != nil {
115
+ return Secp256k1Identity{}, fmt.Errorf("generate secp256k1 private key: %w", err)
116
+ }
117
+ privateKeyHex = hex.EncodeToString(privateKey.Serialize())
118
+ generated = true
119
+ }
120
+
121
+ decoded, normalizedKeyHex, err := decodeSecp256k1PrivateKeyHex(privateKeyHex, true)
122
+ if err != nil {
123
+ return Secp256k1Identity{}, err
124
+ }
125
+
126
+ privateKey := secp256k1.PrivKeyFromBytes(decoded)
127
+ if privateKey == nil {
128
+ return Secp256k1Identity{}, errors.New("invalid secp256k1 private key")
129
+ }
130
+
131
+ publicKeyHex := hex.EncodeToString(privateKey.PubKey().SerializeCompressed())
132
+ address, err := AddressFromCompressedPublicKeyHex(publicKeyHex)
133
+ if err != nil {
134
+ return Secp256k1Identity{}, err
135
+ }
136
+
137
+ return Secp256k1Identity{
138
+ Generated: generated,
139
+ Address: address,
140
+ PublicKey: publicKeyHex,
141
+ PrivateKey: normalizedKeyHex,
142
+ }, nil
143
+}
144
+
145
+func SignSHA256Secp256k1DER(payload []byte, privateKeyHex string) (string, error) {
146
+ decoded, _, err := decodeSecp256k1PrivateKeyHex(privateKeyHex, false)
147
+ if err != nil {
148
+ return "", err
149
+ }
150
+
151
+ hash := sha256.Sum256(payload)
152
+ privateKey := secp256k1.PrivKeyFromBytes(decoded)
153
+ signature := secp256k1ecdsa.Sign(privateKey, hash[:])
154
+ return hex.EncodeToString(signature.Serialize()), nil
155
+}
156
+
157
+func VerifySHA256Secp256k1DER(payload []byte, publicKeyHex, signatureHex string) error {
158
+ pubKeyText := strings.TrimSpace(publicKeyHex)
159
+ if pubKeyText == "" {
160
+ return errors.New("public key is required")
161
+ }
162
+ if strings.HasPrefix(strings.ToLower(pubKeyText), "0x") {
163
+ pubKeyText = pubKeyText[2:]
164
+ }
165
+
166
+ pubKeyBytes, err := hex.DecodeString(pubKeyText)
167
+ if err != nil {
168
+ return errors.New("public key must be hex encoded")
169
+ }
170
+ pubKey, err := secp256k1.ParsePubKey(pubKeyBytes)
171
+ if err != nil {
172
+ return errors.New("invalid secp256k1 public key")
173
+ }
174
+
175
+ sigText := strings.TrimSpace(signatureHex)
176
+ if sigText == "" {
177
+ return errors.New("signature is required")
178
+ }
179
+ if strings.HasPrefix(strings.ToLower(sigText), "0x") {
180
+ sigText = sigText[2:]
181
+ }
182
+
183
+ sigBytes, err := hex.DecodeString(sigText)
184
+ if err != nil {
185
+ return errors.New("signature must be hex encoded")
186
+ }
187
+ signature, err := secp256k1ecdsa.ParseDERSignature(sigBytes)
188
+ if err != nil {
189
+ return fmt.Errorf("parse signature: %w", err)
190
+ }
191
+
192
+ hash := sha256.Sum256(payload)
193
+ if !signature.Verify(hash[:], pubKey) {
194
+ return errors.New("signature is invalid")
195
+ }
196
+ return nil
197
+}
198
+
199
+func NormalizeWireGuardPrivateKey(raw string) (string, error) {
200
+ key, err := decodeWireGuardKey(raw)
201
+ if err != nil {
202
+ return "", err
203
+ }
204
+ clampWireGuardPrivateKey(&key)
205
+ return base64.StdEncoding.EncodeToString(key[:]), nil
206
+}
207
+
208
+func WireGuardPublicKeyFromPrivate(raw string) (string, error) {
209
+ privateKey, err := decodeWireGuardKey(raw)
210
+ if err != nil {
211
+ return "", err
212
+ }
213
+ clampWireGuardPrivateKey(&privateKey)
214
+ var publicKey [32]byte
215
+ curve25519.ScalarBaseMult(&publicKey, &privateKey)
216
+ return base64.StdEncoding.EncodeToString(publicKey[:]), nil
217
+}
218
+
219
+func ValidateWireGuardPublicKey(raw string) error {
220
+ key := strings.TrimSpace(raw)
221
+ if key == "" {
222
+ return errors.New("wireguard_public_key is required")
223
+ }
224
+ decoded, err := base64.StdEncoding.DecodeString(key)
225
+ if err != nil {
226
+ return errors.New("wireguard_public_key must be base64 encoded")
227
+ }
228
+ if len(decoded) != 32 {
229
+ return errors.New("wireguard_public_key must be 32 bytes")
230
+ }
231
+ return nil
232
+}
233
+
234
+func ValidateWireGuardEndpoint(raw string) error {
235
+ endpoint := strings.TrimSpace(raw)
236
+ if endpoint == "" {
237
+ return errors.New("wireguard_endpoint is required")
238
+ }
239
+ host, port, err := net.SplitHostPort(endpoint)
240
+ if err != nil {
241
+ return errors.New("wireguard_endpoint must be host:port")
242
+ }
243
+ if strings.TrimSpace(host) == "" {
244
+ return errors.New("wireguard_endpoint host is required")
245
+ }
246
+ portNum, err := strconv.Atoi(port)
247
+ if err != nil || portNum <= 0 || portNum > 65535 {
248
+ return errors.New("wireguard_endpoint port is invalid")
249
+ }
250
+ return nil
251
+}
252
+
253
+func ValidateOverlayIPv4(raw string) error {
254
+ ipText := strings.TrimSpace(raw)
255
+ if ipText == "" {
256
+ return errors.New("overlay_ipv4 is required")
257
+ }
258
+ ip := net.ParseIP(ipText)
259
+ if ip == nil || ip.To4() == nil {
260
+ return errors.New("overlay_ipv4 must be a valid IPv4 address")
261
+ }
262
+ return nil
263
+}
264
+
265
+func NormalizeOverlayCIDRs(inputs []string) ([]string, error) {
266
+ if len(inputs) == 0 {
267
+ return nil, nil
268
+ }
269
+ seen := make(map[string]struct{}, len(inputs))
270
+ out := make([]string, 0, len(inputs))
271
+ for _, input := range inputs {
272
+ input = strings.TrimSpace(input)
273
+ if input == "" {
274
+ continue
275
+ }
276
+ _, network, err := net.ParseCIDR(input)
277
+ if err != nil {
278
+ return nil, fmt.Errorf("invalid overlay cidr %q", input)
279
+ }
280
+ normalized := network.String()
281
+ if _, ok := seen[normalized]; ok {
282
+ continue
283
+ }
284
+ seen[normalized] = struct{}{}
285
+ out = append(out, normalized)
286
+ }
287
+ sort.Strings(out)
288
+ return out, nil
289
+}
290
+
291
+func WireGuardListenPort(rawEndpoint string) (int, error) {
292
+ endpoint := strings.TrimSpace(rawEndpoint)
293
+ if endpoint == "" {
294
+ return 0, errors.New("wireguard endpoint is required")
295
+ }
296
+ _, portText, err := net.SplitHostPort(endpoint)
297
+ if err != nil {
298
+ return 0, errors.New("wireguard endpoint must be host:port")
299
+ }
300
+ port, err := strconv.Atoi(portText)
301
+ if err != nil || port <= 0 || port > 65535 {
302
+ return 0, errors.New("wireguard endpoint port is invalid")
303
+ }
304
+ return port, nil
305
+}
306
+
307
+func DeriveWireGuardOverlayIPv4(publicKey string) (string, error) {
308
+ decoded, err := base64.StdEncoding.DecodeString(strings.TrimSpace(publicKey))
309
+ if err != nil {
310
+ return "", errors.New("wireguard public key must be base64 encoded")
311
+ }
312
+ if len(decoded) != 32 {
313
+ return "", errors.New("wireguard public key must be 32 bytes")
314
+ }
315
+
316
+ sum := sha256.Sum256(decoded)
317
+ return netip.AddrFrom4([4]byte{
318
+ 100,
319
+ 64 + (sum[0] & 0x3f),
320
+ sum[1],
321
+ 1 + (sum[2] % 254),
322
+ }).String(), nil
323
+}
324
+
325
+func WireGuardKeyHex(raw string) (string, error) {
326
+ key, err := decodeWireGuardKey(raw)
327
+ if err != nil {
328
+ return "", err
329
+ }
330
+ return hex.EncodeToString(key[:]), nil
331
+}
332
+
333
+func decodeSecp256k1PrivateKeyHex(raw string, requireNonZero bool) ([]byte, string, error) {
334
+ privateKeyHex := strings.TrimSpace(raw)
335
+ if privateKeyHex == "" {
336
+ return nil, "", errors.New("private key is required")
337
+ }
338
+ if strings.HasPrefix(strings.ToLower(privateKeyHex), "0x") {
339
+ privateKeyHex = privateKeyHex[2:]
340
+ }
341
+
342
+ decoded, err := hex.DecodeString(privateKeyHex)
343
+ if err != nil {
344
+ return nil, "", errors.New("secp256k1 private key must be hex encoded")
345
+ }
346
+ if len(decoded) != secp256k1.PrivKeyBytesLen {
347
+ return nil, "", fmt.Errorf("secp256k1 private key must be %d bytes", secp256k1.PrivKeyBytesLen)
348
+ }
349
+ if !requireNonZero {
350
+ return decoded, privateKeyHex, nil
351
+ }
352
+
353
+ isZero := true
354
+ for _, b := range decoded {
355
+ if b != 0 {
356
+ isZero = false
357
+ break
358
+ }
359
+ }
360
+ if isZero {
361
+ return nil, "", errors.New("secp256k1 private key must not be zero")
362
+ }
363
+ return decoded, privateKeyHex, nil
364
+}
365
+
366
+func decodeWireGuardKey(raw string) ([32]byte, error) {
367
+ var key [32]byte
368
+ value := strings.TrimSpace(raw)
369
+ if value == "" {
370
+ return key, errors.New("wireguard key is required")
371
+ }
372
+
373
+ var decoded []byte
374
+ var err error
375
+ if len(value) == 64 && !strings.Contains(value, "=") {
376
+ decoded, err = hex.DecodeString(value)
377
+ } else {
378
+ decoded, err = base64.StdEncoding.DecodeString(value)
379
+ }
380
+ if err != nil {
381
+ return key, errors.New("wireguard key must be base64 or hex encoded")
382
+ }
383
+ if len(decoded) != len(key) {
384
+ return key, errors.New("wireguard key must be 32 bytes")
385
+ }
386
+ copy(key[:], decoded)
387
+ return key, nil
388
+}
389
+
390
+func clampWireGuardPrivateKey(key *[32]byte) {
391
+ key[0] &= 248
392
+ key[31] = (key[31] & 127) | 64
393
+}
utils/wireguard.go
deleted
-105
@@ -1,105 +0,0 @@
1
-package utils
2
-
3
-import (
4
- "crypto/sha256"
5
- "encoding/base64"
6
- "encoding/hex"
7
- "errors"
8
- "net"
9
- "net/netip"
10
- "strconv"
11
- "strings"
12
-
13
- "golang.org/x/crypto/curve25519"
14
-)
15
-
16
-func NormalizeWireGuardPrivateKey(raw string) (string, error) {
17
- key, err := decodeWireGuardKey(raw)
18
- if err != nil {
19
- return "", err
20
- }
21
- clampWireGuardPrivateKey(&key)
22
- return base64.StdEncoding.EncodeToString(key[:]), nil
23
-}
24
-
25
-func WireGuardPublicKeyFromPrivate(raw string) (string, error) {
26
- privateKey, err := decodeWireGuardKey(raw)
27
- if err != nil {
28
- return "", err
29
- }
30
- clampWireGuardPrivateKey(&privateKey)
31
- var publicKey [32]byte
32
- curve25519.ScalarBaseMult(&publicKey, &privateKey)
33
- return base64.StdEncoding.EncodeToString(publicKey[:]), nil
34
-}
35
-
36
-func WireGuardListenPort(rawEndpoint string) (int, error) {
37
- endpoint := strings.TrimSpace(rawEndpoint)
38
- if endpoint == "" {
39
- return 0, errors.New("wireguard endpoint is required")
40
- }
41
- _, portText, err := net.SplitHostPort(endpoint)
42
- if err != nil {
43
- return 0, errors.New("wireguard endpoint must be host:port")
44
- }
45
- port, err := strconv.Atoi(portText)
46
- if err != nil || port <= 0 || port > 65535 {
47
- return 0, errors.New("wireguard endpoint port is invalid")
48
- }
49
- return port, nil
50
-}
51
-
52
-func DeriveWireGuardOverlayIPv4(publicKey string) (string, error) {
53
- decoded, err := base64.StdEncoding.DecodeString(strings.TrimSpace(publicKey))
54
- if err != nil {
55
- return "", errors.New("wireguard public key must be base64 encoded")
56
- }
57
- if len(decoded) != 32 {
58
- return "", errors.New("wireguard public key must be 32 bytes")
59
- }
60
-
61
- sum := sha256.Sum256(decoded)
62
- return netip.AddrFrom4([4]byte{
63
- 100,
64
- 64 + (sum[0] & 0x3f),
65
- sum[1],
66
- 1 + (sum[2] % 254),
67
- }).String(), nil
68
-}
69
-
70
-func WireGuardKeyHex(raw string) (string, error) {
71
- key, err := decodeWireGuardKey(raw)
72
- if err != nil {
73
- return "", err
74
- }
75
- return hex.EncodeToString(key[:]), nil
76
-}
77
-
78
-func decodeWireGuardKey(raw string) ([32]byte, error) {
79
- var key [32]byte
80
- value := strings.TrimSpace(raw)
81
- if value == "" {
82
- return key, errors.New("wireguard key is required")
83
- }
84
-
85
- var decoded []byte
86
- var err error
87
- if len(value) == 64 && !strings.Contains(value, "=") {
88
- decoded, err = hex.DecodeString(value)
89
- } else {
90
- decoded, err = base64.StdEncoding.DecodeString(value)
91
- }
92
- if err != nil {
93
- return key, errors.New("wireguard key must be base64 or hex encoded")
94
- }
95
- if len(decoded) != len(key) {
96
- return key, errors.New("wireguard key must be 32 bytes")
97
- }
98
- copy(key[:], decoded)
99
- return key, nil
100
-}
101
-
102
-func clampWireGuardPrivateKey(key *[32]byte) {
103
- key[0] &= 248
104
- key[31] = (key[31] & 127) | 64
105
-}