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