feat: add relay discovery self-announce

Kim committed Apr 14, 2026 at 16:33 UTC 94fca4004d1040d49c5b512f2348fb81ee363c43
22 files changed +751 -1347
docs/src/routes/api-reference/+page.md
+23 -3
@@ -105,6 +105,7 @@ Admin clients authenticate using a shared secret key:
105 |--------|------|-------------|------|
106 | `GET` | `/healthz` | Health check | None |
107 | `GET` | `/discovery` | Relay discovery | None |
108 +| `POST` | `/discovery/announce` | Relay discovery self-announce | Signed Descriptor |
109 | `GET` | `/v1/sign` | Keyless TLS signing | None |
110 | `GET` | `/thumbnail/{hostname}` | Cached thumbnail screenshot | None |
111 | `GET` | `/tunnel/status` | Tunnel connection status | Access Token |
@@ -130,7 +131,7 @@ Returns relay health status.
131
132 ### `GET /discovery`
133
133 -Returns relay discovery information including the relay's own descriptor and any known peer relays. Only available when discovery is enabled in the server configuration.
134 +Returns signed relay discovery descriptors for this relay and any known peer relays. Only available when discovery is enabled in the server configuration.
135
136 **Response fields:**
137
@@ -138,8 +139,7 @@ Returns relay discovery information including the relay's own descriptor and any
139 |-------|------|-------------|
140 | `protocol_version` | `string` | Protocol version identifier |
141 | `generated_at` | `string` | ISO 8601 timestamp |
141 -| `self` | `RelayDescriptor` | This relay's descriptor |
142 -| `relays` | `RelayDescriptor[]` | Known peer relay descriptors |
142 +| `relays` | `RelayDescriptor[]` | Signed descriptors for this relay and known peer relays |
143
144 **Example:**
145
@@ -147,6 +147,26 @@ Returns relay discovery information including the relay's own descriptor and any
147 curl https://relay.example.com/discovery
148 ```
149
150 +### `POST /discovery/announce`
151 +
152 +Submits this relay's signed descriptor to a bootstrap relay so registry-external relays can enter the discovery mesh. Relays self-announce periodically when discovery is enabled.
153 +
154 +**Auth:** Signed relay descriptor
155 +
156 +**Request fields:**
157 +
158 +| Field | Type | Required | Description |
159 +|-------|------|----------|-------------|
160 +| `protocol_version` | `string` | No | Discovery protocol version |
161 +| `descriptor` | `RelayDescriptor` | Yes | Signed relay descriptor |
162 +
163 +**Response fields:**
164 +
165 +| Field | Type | Description |
166 +|-------|------|-------------|
167 +| `protocol_version` | `string` | Discovery protocol version |
168 +| `accepted` | `boolean` | Whether the descriptor was accepted |
169 +
170 ### `GET /v1/sign`
171
172 Keyless TLS signing endpoint. Used by the relay's keyless TLS infrastructure. Only available when the API server is configured with a TLS private key.
docs/src/routes/architecture/+page.md
+2 -2
@@ -292,8 +292,8 @@ Result: raw public UDP exposure with an internal QUIC datagram backhaul. UDP and
292
293 ## WireGuard Overlay and Discovery
294
295 -- Discovery bootstraps from public HTTPS relay URLs, then optionally synchronizes over WireGuard overlay.
296 -- Discovery descriptors are transport-authenticated by the queried relay endpoint, not by embedded signatures. Independent `domain -> address` verification comes from optional ENS/DNSSEC evidence, not from the discovery payload itself.
295 +- Discovery bootstraps from public HTTPS relay URLs, then expands through relay-to-relay `/discovery` polling and periodic self-announces to bootstrap relays through `/discovery/announce`.
296 +- Discovery descriptors carry secp256k1 signatures that bind relay routing metadata to the relay identity. Lease access tokens remain separate and authorize tenant lease operations only.
297 - The overlay peer API is plain HTTP on the WireGuard network, not public Internet HTTP. It serves the same discovery payload shape used by public `/discovery`.
298 - Overlay failure affects inter-relay discovery and mesh synchronization only. Tenant stream routing, keyless TLS, register/renew/connect, and public UDP ingress do not depend on the WireGuard transport path.
299
portal/api_server.go
+32 -14
@@ -157,8 +157,12 @@ func (s *Server) extractAllowedClientIP(w http.ResponseWriter, r *http.Request)
157 return "", false
158 }
159
160 -func (s *Server) handleRelayDiscovery(w http.ResponseWriter, r *http.Request) {
161 - now := time.Now().UTC()
160 +func (s *Server) signedRelayDescriptor(now time.Time) (types.RelayDescriptor, error) {
161 + if now.IsZero() {
162 + now = time.Now().UTC()
163 + } else {
164 + now = now.UTC()
165 + }
166 activeConns := float64(s.proxy.ActiveConns())
167 tcpTrafficBPS := s.proxy.CurrentTCPBPS(now)
168 ingressAddr := s.identity.Name
@@ -176,7 +180,7 @@ func (s *Server) handleRelayDiscovery(w http.ResponseWriter, r *http.Request) {
180 overlayCIDRs = append([]string(nil), cfg.OverlayCIDRs...)
181 }
182
179 - self, err := utils.NormalizeDescriptor(types.RelayDescriptor{
183 + self := types.RelayDescriptor{
184 Identity: s.identity.Base(),
185 RelayID: s.cfg.PortalURL,
186 OwnerAddress: s.identity.Address,
@@ -196,17 +200,36 @@ func (s *Server) handleRelayDiscovery(w http.ResponseWriter, r *http.Request) {
200 Load: activeConns,
201 LoadScore: tcpTrafficBPS,
202 LastUpdated: now.UnixMilli(),
199 - })
203 + }
204 +
205 + signedSelf, err := auth.SignRelayDescriptor(self, s.identity.PrivateKey)
206 if err != nil {
201 - utils.WriteAPIError(w, http.StatusInternalServerError, types.APIErrorCodeInternal, err.Error())
207 + return types.RelayDescriptor{}, err
208 + }
209 + return signedSelf, nil
210 +}
211 +
212 +func (s *Server) handleRelayDiscovery(w http.ResponseWriter, r *http.Request) {
213 + if !utils.RequireMethod(w, r, http.MethodGet) {
214 return
215 }
204 - signedSelf, err := discovery.SignDescriptor(self, s.identity.PrivateKey)
216 + if s.relaySet == nil {
217 + utils.WriteAPIError(w, http.StatusServiceUnavailable, types.APIErrorCodeFeatureUnavailable, "relay discovery disabled")
218 + return
219 + }
220 +
221 + now := time.Now().UTC()
222 + self, err := s.signedRelayDescriptor(now)
223 if err != nil {
224 utils.WriteAPIError(w, http.StatusInternalServerError, types.APIErrorCodeInternal, err.Error())
225 return
226 }
209 - s.relaySet.ServeDiscovery(w, r, signedSelf)
227 +
228 + utils.WriteAPIData(w, http.StatusOK, types.DiscoveryResponse{
229 + ProtocolVersion: types.DiscoveryVersion,
230 + GeneratedAt: now,
231 + Relays: s.relaySet.Descriptors(self),
232 + })
233 }
234
235 func (s *Server) handleRelayDiscoveryAnnounce(w http.ResponseWriter, r *http.Request) {
@@ -249,15 +272,10 @@ func (s *Server) handleRelayDiscoveryAnnounce(w http.ResponseWriter, r *http.Req
272 }
273
274 now := time.Now().UTC()
252 - accepted, _, err := s.relaySet.InsertAnnounced(desc, now)
253 - if err != nil {
275 + if err := s.relaySet.InsertAnnounced(desc, now); err != nil {
276 utils.WriteAPIError(w, http.StatusBadRequest, types.APIErrorCodeInvalidRequest, err.Error())
277 return
278 }
257 - if !accepted {
258 - utils.WriteAPIError(w, http.StatusConflict, types.APIErrorCodeInvalidRequest, "announce not accepted")
259 - return
260 - }
279
280 log.Info().
281 Str("relay", desc.APIHTTPSAddr).
@@ -301,7 +319,7 @@ func (s *Server) handleRegister(w http.ResponseWriter, r *http.Request) {
319 challenge, err := s.registry.consumeVerifiedRegisterChallenge(req)
320 if err != nil {
321 switch {
304 - case errors.Is(err, auth.ErrInvalidSignature):
322 + case errors.Is(err, auth.ErrRegisterChallengeInvalidSignature):
323 utils.WriteAPIError(w, http.StatusForbidden, types.APIErrorCodeUnauthorized, err.Error())
324 default:
325 utils.InvalidRequestError(err).Write(w)
portal/auth/auth.go deleted
-269
@@ -1,269 +0,0 @@
1 -package auth
2 -
3 -import (
4 - "crypto/sha256"
5 - "errors"
6 - "fmt"
7 - "strings"
8 - "time"
9 -
10 - "github.com/decred/dcrd/dcrec/secp256k1/v4"
11 - "github.com/decred/dcrd/dcrec/secp256k1/v4/ecdsa"
12 - jose "github.com/go-jose/go-jose/v4"
13 - "github.com/go-jose/go-jose/v4/jwt"
14 - "github.com/spruceid/siwe-go"
15 -
16 - "github.com/gosuda/portal-tunnel/v2/types"
17 - "github.com/gosuda/portal-tunnel/v2/utils"
18 -)
19 -
20 -const (
21 - registerStatement = "Register a portal lease"
22 - leaseAccessTokenAudience = "portal-sdk"
23 -)
24 -
25 -var (
26 - ErrChallengeExpired = errors.New("register challenge expired")
27 - ErrChallengeNotFound = errors.New("register challenge not found")
28 - ErrInvalidSignature = errors.New("siwe signature is invalid")
29 - ErrMessageMismatch = errors.New("siwe message does not match register challenge")
30 -)
31 -
32 -const leaseTokenAlgorithm = jose.SignatureAlgorithm("ES256K")
33 -
34 -type LeaseAccessTokenClaims struct {
35 - jwt.Claims
36 - Identity types.Identity `json:"identity"`
37 -}
38 -
39 -type es256kOpaqueSigner struct {
40 - keyID string
41 - privateKey *secp256k1.PrivateKey
42 -}
43 -
44 -func (s *es256kOpaqueSigner) Public() *jose.JSONWebKey {
45 - return &jose.JSONWebKey{KeyID: s.keyID}
46 -}
47 -
48 -func (s *es256kOpaqueSigner) Algs() []jose.SignatureAlgorithm {
49 - return []jose.SignatureAlgorithm{leaseTokenAlgorithm}
50 -}
51 -
52 -func (s *es256kOpaqueSigner) SignPayload(payload []byte, alg jose.SignatureAlgorithm) ([]byte, error) {
53 - if alg != leaseTokenAlgorithm {
54 - return nil, jose.ErrUnsupportedAlgorithm
55 - }
56 - if s == nil || s.privateKey == nil {
57 - return nil, errors.New("signing key is required")
58 - }
59 -
60 - hash := sha256.Sum256(payload)
61 - compact := ecdsa.SignCompact(s.privateKey, hash[:], false)
62 - if len(compact) != 65 {
63 - return nil, errors.New("invalid compact signature length")
64 - }
65 -
66 - signature := make([]byte, 64)
67 - copy(signature[:32], compact[1:33])
68 - copy(signature[32:], compact[33:65])
69 - return signature, nil
70 -}
71 -
72 -type es256kOpaqueVerifier struct {
73 - publicKey *secp256k1.PublicKey
74 -}
75 -
76 -func (v *es256kOpaqueVerifier) VerifyPayload(payload []byte, signature []byte, alg jose.SignatureAlgorithm) error {
77 - if alg != leaseTokenAlgorithm {
78 - return jose.ErrUnsupportedAlgorithm
79 - }
80 - if v == nil || v.publicKey == nil {
81 - return errors.New("verification key is required")
82 - }
83 - if len(signature) != 64 {
84 - return errors.New("invalid es256k signature length")
85 - }
86 -
87 - var r, s secp256k1.ModNScalar
88 - if overflow := r.SetByteSlice(signature[:32]); overflow || r.IsZero() {
89 - return errors.New("invalid es256k signature r")
90 - }
91 - if overflow := s.SetByteSlice(signature[32:]); overflow || s.IsZero() {
92 - return errors.New("invalid es256k signature s")
93 - }
94 -
95 - hash := sha256.Sum256(payload)
96 - return verifyRawSignature(hash[:], &r, &s, v.publicKey)
97 -}
98 -
99 -func verifyRawSignature(hash []byte, r, s *secp256k1.ModNScalar, publicKey *secp256k1.PublicKey) error {
100 - signature := ecdsa.NewSignature(r, s)
101 - if !signature.Verify(hash, publicKey) {
102 - return errors.New("token signature is invalid")
103 - }
104 - return nil
105 -}
106 -
107 -type RegisterChallenge struct {
108 - ChallengeID string
109 - ExpiresAt time.Time
110 - Request types.RegisterChallengeRequest
111 - SIWEMessage string
112 -
113 - domain string
114 - nonce string
115 -}
116 -
117 -func NewRegisterChallenge(req types.RegisterChallengeRequest, domain, uri string, now time.Time, ttl time.Duration) (*RegisterChallenge, error) {
118 - normalizedIdentity, err := utils.NormalizeIdentity(req.Identity)
119 - if err != nil {
120 - return nil, err
121 - }
122 -
123 - challengeID := utils.RandomID("rch_")
124 - nonce := siwe.GenerateNonce()
125 - expiresAt := now.UTC().Add(ttl)
126 - siweMessage, err := BuildRegisterChallengeMessage(domain, normalizedIdentity.Address, uri, challengeID, nonce, now.UTC(), expiresAt)
127 - if err != nil {
128 - return nil, err
129 - }
130 -
131 - normalizedRequest := types.RegisterChallengeRequest{
132 - Identity: normalizedIdentity,
133 - Metadata: req.Metadata.Copy(),
134 - TTL: req.TTL,
135 - UDPEnabled: req.UDPEnabled,
136 - TCPEnabled: req.TCPEnabled,
137 - }
138 -
139 - return &RegisterChallenge{
140 - ChallengeID: challengeID,
141 - ExpiresAt: expiresAt,
142 - Request: normalizedRequest,
143 - SIWEMessage: siweMessage,
144 - domain: strings.TrimSpace(domain),
145 - nonce: nonce,
146 - }, nil
147 -}
148 -
149 -func BuildRegisterChallengeMessage(domain, address, uri, challengeID, nonce string, issuedAt, expiresAt time.Time) (string, error) {
150 - message, err := siwe.InitMessage(domain, address, uri, nonce, map[string]interface{}{
151 - "statement": registerStatement,
152 - "chainId": 1,
153 - "issuedAt": issuedAt.UTC().Format(time.RFC3339),
154 - "expirationTime": expiresAt.UTC().Format(time.RFC3339),
155 - "requestId": challengeID,
156 - })
157 - if err != nil {
158 - return "", fmt.Errorf("build siwe message: %w", err)
159 - }
160 - return message.String(), nil
161 -}
162 -
163 -func (c *RegisterChallenge) Expired(now time.Time) bool {
164 - if c == nil {
165 - return true
166 - }
167 - return now.UTC().After(c.ExpiresAt)
168 -}
169 -
170 -func (c *RegisterChallenge) Verify(req types.RegisterRequest, now time.Time) error {
171 - if c == nil {
172 - return ErrChallengeNotFound
173 - }
174 - if strings.TrimSpace(req.SIWEMessage) != c.SIWEMessage {
175 - return ErrMessageMismatch
176 - }
177 - if err := VerifyRegisterChallengeMessage(c.SIWEMessage, req.SIWESignature, c.domain, c.nonce, now.UTC()); err != nil {
178 - return ErrInvalidSignature
179 - }
180 - return nil
181 -}
182 -
183 -func VerifyRegisterChallengeMessage(messageText, signature, domain, nonce string, now time.Time) error {
184 - message, err := siwe.ParseMessage(strings.TrimSpace(messageText))
185 - if err != nil {
186 - return err
187 - }
188 - normalizedDomain := strings.TrimSpace(domain)
189 - normalizedNonce := strings.TrimSpace(nonce)
190 - verifiedAt := now.UTC()
191 - _, err = message.Verify(strings.TrimSpace(signature), &normalizedDomain, &normalizedNonce, &verifiedAt)
192 - return err
193 -}
194 -
195 -func IssueLeaseAccessToken(privateKeyHex, keyID, issuer string, identity types.Identity, ttl time.Duration) (string, LeaseAccessTokenClaims, error) {
196 - privateKey, _, err := utils.ParseSecp256k1PrivateKeyHex(privateKeyHex, false)
197 - if err != nil {
198 - return "", LeaseAccessTokenClaims{}, err
199 - }
200 - normalizedIdentity, err := utils.NormalizeIdentity(identity)
201 - if err != nil {
202 - return "", LeaseAccessTokenClaims{}, err
203 - }
204 -
205 - signer, err := jose.NewSigner(jose.SigningKey{
206 - Algorithm: leaseTokenAlgorithm,
207 - Key: &es256kOpaqueSigner{
208 - keyID: strings.TrimSpace(keyID),
209 - privateKey: privateKey,
210 - },
211 - }, (&jose.SignerOptions{}).WithType("JWT"))
212 - if err != nil {
213 - return "", LeaseAccessTokenClaims{}, err
214 - }
215 -
216 - now := time.Now().UTC()
217 - expiresAt := now.Add(ttl)
218 - claims := LeaseAccessTokenClaims{
219 - Claims: jwt.Claims{
220 - Issuer: strings.TrimSpace(issuer),
221 - Subject: normalizedIdentity.Key(),
222 - Audience: jwt.Audience{leaseAccessTokenAudience},
223 - ID: utils.RandomID("tok_"),
224 - IssuedAt: jwt.NewNumericDate(now),
225 - NotBefore: jwt.NewNumericDate(now),
226 - Expiry: jwt.NewNumericDate(expiresAt),
227 - },
228 - Identity: normalizedIdentity,
229 - }
230 -
231 - token, err := jwt.Signed(signer).Claims(claims).Serialize()
232 - if err != nil {
233 - return "", LeaseAccessTokenClaims{}, err
234 - }
235 - return token, claims, nil
236 -}
237 -
238 -func VerifyLeaseAccessToken(token, publicKeyHex, issuer string, now time.Time) (LeaseAccessTokenClaims, error) {
239 - publicKey, err := utils.ParseSecp256k1PublicKeyHex(publicKeyHex)
240 - if err != nil {
241 - return LeaseAccessTokenClaims{}, err
242 - }
243 -
244 - parsed, err := jwt.ParseSigned(strings.TrimSpace(token), []jose.SignatureAlgorithm{leaseTokenAlgorithm})
245 - if err != nil {
246 - return LeaseAccessTokenClaims{}, err
247 - }
248 -
249 - var claims LeaseAccessTokenClaims
250 - if err := parsed.Claims(&es256kOpaqueVerifier{publicKey: publicKey}, &claims); err != nil {
251 - return LeaseAccessTokenClaims{}, err
252 - }
253 - normalizedClaimsIdentity, err := utils.NormalizeIdentity(claims.Identity)
254 - if err != nil {
255 - return LeaseAccessTokenClaims{}, err
256 - }
257 - if normalizedClaimsIdentity.Key() != claims.Subject {
258 - return LeaseAccessTokenClaims{}, errors.New("lease access token identity does not match subject")
259 - }
260 - claims.Identity = normalizedClaimsIdentity
261 - if err := claims.ValidateWithLeeway(jwt.Expected{
262 - Issuer: strings.TrimSpace(issuer),
263 - AnyAudience: jwt.Audience{leaseAccessTokenAudience},
264 - Time: now.UTC(),
265 - }, 0); err != nil {
266 - return LeaseAccessTokenClaims{}, err
267 - }
268 - return claims, nil
269 -}
portal/auth/lease_token.go new
+143
@@ -0,0 +1,143 @@
1 +package auth
2 +
3 +import (
4 + "errors"
5 + "strings"
6 + "time"
7 +
8 + "github.com/decred/dcrd/dcrec/secp256k1/v4"
9 + jose "github.com/go-jose/go-jose/v4"
10 + "github.com/go-jose/go-jose/v4/jwt"
11 +
12 + "github.com/gosuda/portal-tunnel/v2/types"
13 + "github.com/gosuda/portal-tunnel/v2/utils"
14 +)
15 +
16 +const (
17 + leaseAccessTokenAudience = "portal-sdk"
18 + leaseTokenAlgorithm = jose.SignatureAlgorithm("ES256K")
19 +)
20 +
21 +type LeaseAccessTokenClaims struct {
22 + jwt.Claims
23 + Identity types.Identity `json:"identity"`
24 +}
25 +
26 +type es256kOpaqueSigner struct {
27 + keyID string
28 + privateKey *secp256k1.PrivateKey
29 +}
30 +
31 +func (s *es256kOpaqueSigner) Public() *jose.JSONWebKey {
32 + return &jose.JSONWebKey{KeyID: s.keyID}
33 +}
34 +
35 +func (s *es256kOpaqueSigner) Algs() []jose.SignatureAlgorithm {
36 + return []jose.SignatureAlgorithm{leaseTokenAlgorithm}
37 +}
38 +
39 +func (s *es256kOpaqueSigner) SignPayload(payload []byte, alg jose.SignatureAlgorithm) ([]byte, error) {
40 + if alg != leaseTokenAlgorithm {
41 + return nil, jose.ErrUnsupportedAlgorithm
42 + }
43 + if s == nil || s.privateKey == nil {
44 + return nil, errors.New("signing key is required")
45 + }
46 + return utils.SignSHA256Secp256k1Raw64(payload, s.privateKey)
47 +}
48 +
49 +type es256kOpaqueVerifier struct {
50 + publicKey *secp256k1.PublicKey
51 +}
52 +
53 +func (v *es256kOpaqueVerifier) VerifyPayload(payload []byte, signature []byte, alg jose.SignatureAlgorithm) error {
54 + if alg != leaseTokenAlgorithm {
55 + return jose.ErrUnsupportedAlgorithm
56 + }
57 + if v == nil || v.publicKey == nil {
58 + return errors.New("verification key is required")
59 + }
60 + if err := utils.VerifySHA256Secp256k1Raw64(payload, signature, v.publicKey); err != nil {
61 + if errors.Is(err, utils.ErrSecp256k1SignatureInvalid) {
62 + return errors.New("token signature is invalid")
63 + }
64 + return err
65 + }
66 + return nil
67 +}
68 +
69 +func IssueLeaseAccessToken(privateKeyHex, keyID, issuer string, identity types.Identity, ttl time.Duration) (string, LeaseAccessTokenClaims, error) {
70 + privateKey, _, err := utils.ParseSecp256k1PrivateKeyHex(privateKeyHex, false)
71 + if err != nil {
72 + return "", LeaseAccessTokenClaims{}, err
73 + }
74 + normalizedIdentity, err := utils.NormalizeIdentity(identity)
75 + if err != nil {
76 + return "", LeaseAccessTokenClaims{}, err
77 + }
78 +
79 + signer, err := jose.NewSigner(jose.SigningKey{
80 + Algorithm: leaseTokenAlgorithm,
81 + Key: &es256kOpaqueSigner{
82 + keyID: strings.TrimSpace(keyID),
83 + privateKey: privateKey,
84 + },
85 + }, (&jose.SignerOptions{}).WithType("JWT"))
86 + if err != nil {
87 + return "", LeaseAccessTokenClaims{}, err
88 + }
89 +
90 + now := time.Now().UTC()
91 + expiresAt := now.Add(ttl)
92 + claims := LeaseAccessTokenClaims{
93 + Claims: jwt.Claims{
94 + Issuer: strings.TrimSpace(issuer),
95 + Subject: normalizedIdentity.Key(),
96 + Audience: jwt.Audience{leaseAccessTokenAudience},
97 + ID: utils.RandomID("tok_"),
98 + IssuedAt: jwt.NewNumericDate(now),
99 + NotBefore: jwt.NewNumericDate(now),
100 + Expiry: jwt.NewNumericDate(expiresAt),
101 + },
102 + Identity: normalizedIdentity,
103 + }
104 +
105 + token, err := jwt.Signed(signer).Claims(claims).Serialize()
106 + if err != nil {
107 + return "", LeaseAccessTokenClaims{}, err
108 + }
109 + return token, claims, nil
110 +}
111 +
112 +func VerifyLeaseAccessToken(token, publicKeyHex, issuer string, now time.Time) (LeaseAccessTokenClaims, error) {
113 + publicKey, err := utils.ParseSecp256k1PublicKeyHex(publicKeyHex)
114 + if err != nil {
115 + return LeaseAccessTokenClaims{}, err
116 + }
117 +
118 + parsed, err := jwt.ParseSigned(strings.TrimSpace(token), []jose.SignatureAlgorithm{leaseTokenAlgorithm})
119 + if err != nil {
120 + return LeaseAccessTokenClaims{}, err
121 + }
122 +
123 + var claims LeaseAccessTokenClaims
124 + if err := parsed.Claims(&es256kOpaqueVerifier{publicKey: publicKey}, &claims); err != nil {
125 + return LeaseAccessTokenClaims{}, err
126 + }
127 + normalizedClaimsIdentity, err := utils.NormalizeIdentity(claims.Identity)
128 + if err != nil {
129 + return LeaseAccessTokenClaims{}, err
130 + }
131 + if normalizedClaimsIdentity.Key() != claims.Subject {
132 + return LeaseAccessTokenClaims{}, errors.New("lease access token identity does not match subject")
133 + }
134 + claims.Identity = normalizedClaimsIdentity
135 + if err := claims.ValidateWithLeeway(jwt.Expected{
136 + Issuer: strings.TrimSpace(issuer),
137 + AnyAudience: jwt.Audience{leaseAccessTokenAudience},
138 + Time: now.UTC(),
139 + }, 0); err != nil {
140 + return LeaseAccessTokenClaims{}, err
141 + }
142 + return claims, nil
143 +}
portal/auth/register_challenge.go new
+91
@@ -0,0 +1,91 @@
1 +package auth
2 +
3 +import (
4 + "errors"
5 + "fmt"
6 + "strings"
7 + "time"
8 +
9 + "github.com/spruceid/siwe-go"
10 +
11 + "github.com/gosuda/portal-tunnel/v2/types"
12 + "github.com/gosuda/portal-tunnel/v2/utils"
13 +)
14 +
15 +var (
16 + ErrRegisterChallengeExpired = errors.New("register challenge expired")
17 + ErrRegisterChallengeNotFound = errors.New("register challenge not found")
18 + ErrRegisterChallengeInvalidSignature = errors.New("siwe signature is invalid")
19 +)
20 +
21 +type RegisterChallenge struct {
22 + ChallengeID string
23 + ExpiresAt time.Time
24 + Request types.RegisterChallengeRequest
25 + SIWEMessage string
26 +
27 + domain string
28 + nonce string
29 +}
30 +
31 +func NewRegisterChallenge(req types.RegisterChallengeRequest, domain, uri string, now time.Time, ttl time.Duration) (*RegisterChallenge, error) {
32 + normalizedIdentity, err := utils.NormalizeIdentity(req.Identity)
33 + if err != nil {
34 + return nil, err
35 + }
36 +
37 + challengeID := utils.RandomID("rch_")
38 + nonce := siwe.GenerateNonce()
39 + expiresAt := now.UTC().Add(ttl)
40 + message, err := siwe.InitMessage(domain, normalizedIdentity.Address, uri, nonce, map[string]interface{}{
41 + "statement": "Register a portal lease",
42 + "chainId": 1,
43 + "issuedAt": now.UTC().Format(time.RFC3339),
44 + "expirationTime": expiresAt.UTC().Format(time.RFC3339),
45 + "requestId": challengeID,
46 + })
47 + if err != nil {
48 + return nil, fmt.Errorf("build siwe message: %w", err)
49 + }
50 +
51 + normalizedRequest := types.RegisterChallengeRequest{
52 + Identity: normalizedIdentity,
53 + Metadata: req.Metadata.Copy(),
54 + TTL: req.TTL,
55 + UDPEnabled: req.UDPEnabled,
56 + TCPEnabled: req.TCPEnabled,
57 + }
58 +
59 + return &RegisterChallenge{
60 + ChallengeID: challengeID,
61 + ExpiresAt: expiresAt,
62 + Request: normalizedRequest,
63 + SIWEMessage: message.String(),
64 + domain: strings.TrimSpace(domain),
65 + nonce: nonce,
66 + }, nil
67 +}
68 +
69 +func (c *RegisterChallenge) Expired(now time.Time) bool {
70 + return c == nil || now.After(c.ExpiresAt)
71 +}
72 +
73 +func (c *RegisterChallenge) Verify(req types.RegisterRequest, now time.Time) error {
74 + if c == nil {
75 + return ErrRegisterChallengeNotFound
76 + }
77 + if strings.TrimSpace(req.SIWEMessage) != c.SIWEMessage {
78 + return errors.New("siwe message does not match register challenge")
79 + }
80 + message, err := siwe.ParseMessage(strings.TrimSpace(c.SIWEMessage))
81 + if err != nil {
82 + return ErrRegisterChallengeInvalidSignature
83 + }
84 + normalizedDomain := strings.TrimSpace(c.domain)
85 + normalizedNonce := strings.TrimSpace(c.nonce)
86 + verifiedAt := now.UTC()
87 + if _, err := message.Verify(strings.TrimSpace(req.SIWESignature), &normalizedDomain, &normalizedNonce, &verifiedAt); err != nil {
88 + return ErrRegisterChallengeInvalidSignature
89 + }
90 + return nil
91 +}
portal/auth/relay_descriptor.go new
+96
@@ -0,0 +1,96 @@
1 +package auth
2 +
3 +import (
4 + "encoding/base64"
5 + "encoding/hex"
6 + "errors"
7 + "fmt"
8 + "strings"
9 +
10 + "github.com/gosuda/portal-tunnel/v2/types"
11 + "github.com/gosuda/portal-tunnel/v2/utils"
12 +)
13 +
14 +// SignRelayDescriptor returns a copy of desc with its Signature field
15 +// populated by signing the canonical bytes with the supplied secp256k1
16 +// private key (hex encoded). The signature is recoverable, so verifiers do
17 +// not need to know the public key out of band; they recover it from the
18 +// signature and check it derives the descriptor's Address field.
19 +//
20 +// Mutable telemetry fields (Load, LoadScore, LastUpdated) are NOT covered by
21 +// the signature, so callers may update them after signing without
22 +// invalidating the signature.
23 +func SignRelayDescriptor(desc types.RelayDescriptor, privateKeyHex string) (types.RelayDescriptor, error) {
24 + privateKey, _, err := utils.ParseSecp256k1PrivateKeyHex(privateKeyHex, true)
25 + if err != nil {
26 + return types.RelayDescriptor{}, fmt.Errorf("relay descriptor signing key: %w", err)
27 + }
28 +
29 + desc.Signature = ""
30 + normalized, err := utils.NormalizeDescriptor(desc)
31 + if err != nil {
32 + return types.RelayDescriptor{}, fmt.Errorf("normalize relay descriptor for signing: %w", err)
33 + }
34 + if strings.TrimSpace(normalized.Address) == "" {
35 + return types.RelayDescriptor{}, errors.New("relay descriptor address is required for signature verification")
36 + }
37 + desc = normalized
38 +
39 + canonical, err := types.CanonicalBytes(desc)
40 + if err != nil {
41 + return types.RelayDescriptor{}, fmt.Errorf("canonicalize relay descriptor: %w", err)
42 + }
43 + signature, err := utils.SignSHA256Secp256k1Compact(canonical, privateKey, true)
44 + if err != nil {
45 + return types.RelayDescriptor{}, err
46 + }
47 +
48 + desc.Signature = base64.StdEncoding.EncodeToString(signature)
49 + return desc, nil
50 +}
51 +
52 +// VerifyRelayDescriptor checks the descriptor's signature against its
53 +// canonical bytes and confirms that the recovered signing key corresponds to
54 +// the descriptor's Address field. It returns the verified normalized
55 +// descriptor on success.
56 +func VerifyRelayDescriptor(desc types.RelayDescriptor) (types.RelayDescriptor, error) {
57 + rawSignature := strings.TrimSpace(desc.Signature)
58 + if rawSignature == "" {
59 + return types.RelayDescriptor{}, errors.New("relay descriptor is not signed")
60 + }
61 +
62 + signature, err := base64.StdEncoding.DecodeString(rawSignature)
63 + if err != nil {
64 + return types.RelayDescriptor{}, fmt.Errorf("relay descriptor signature is invalid: base64 decode: %w", err)
65 + }
66 +
67 + unsignedCopy := desc
68 + unsignedCopy.Signature = ""
69 + normalized, err := utils.NormalizeDescriptor(unsignedCopy)
70 + if err != nil {
71 + return types.RelayDescriptor{}, fmt.Errorf("relay descriptor signature is invalid: normalize: %w", err)
72 + }
73 + if strings.TrimSpace(normalized.Address) == "" {
74 + return types.RelayDescriptor{}, errors.New("relay descriptor address is required for signature verification")
75 + }
76 + canonical, err := types.CanonicalBytes(normalized)
77 + if err != nil {
78 + return types.RelayDescriptor{}, fmt.Errorf("canonicalize relay descriptor: %w", err)
79 + }
80 +
81 + publicKey, err := utils.RecoverSHA256Secp256k1Compact(canonical, signature)
82 + if err != nil {
83 + return types.RelayDescriptor{}, fmt.Errorf("relay descriptor signature is invalid: %w", err)
84 + }
85 +
86 + publicKeyHex := hex.EncodeToString(publicKey.SerializeCompressed())
87 + derivedAddress, err := utils.AddressFromCompressedPublicKeyHex(publicKeyHex)
88 + if err != nil {
89 + return types.RelayDescriptor{}, fmt.Errorf("derive address from recovered key: %w", err)
90 + }
91 + if !strings.EqualFold(strings.TrimSpace(derivedAddress), strings.TrimSpace(normalized.Address)) {
92 + return types.RelayDescriptor{}, errors.New("relay descriptor address does not match recovered signing key")
93 + }
94 + normalized.Signature = rawSignature
95 + return normalized, nil
96 +}
portal/discovery/announce_adversarial_test.go deleted
-122
@@ -1,122 +0,0 @@
1 -package discovery
2 -
3 -import (
4 - "testing"
5 - "time"
6 -
7 - "github.com/gosuda/portal-tunnel/v2/types"
8 -)
9 -
10 -// TestLRUEvictionLeaksRollbackHistoryForSameIdentity exercises a replay-via-LRU
11 -// attack. The rollback defense in upsertDescriptorLocked is anchored in
12 -// s.keyIndex, which records the newest IssuedAt ever accepted for a signing
13 -// identity. deleteRelayLocked drops the keyIndex entry as soon as the last
14 -// URL slot for that identity is evicted — so if the legitimate relay's slot
15 -// is pushed out by LRU pressure, the rollback history is lost.
16 -//
17 -// An attacker who captured an older (but still unexpired) signed descriptor
18 -// can then re-announce it and the receiving relay will accept it as if it had
19 -// never seen any newer IssuedAt. That is strictly contrary to the
20 -// "monotonic IssuedAt per signing identity" invariant documented at
21 -// relayset.go on upsertDescriptorLocked.
22 -//
23 -// This test must FAIL on a correct implementation (i.e. InsertAnnounced must
24 -// reject the older descriptor), and PASS (in the failing-assertion sense)
25 -// on the current implementation — surfacing the bug.
26 -func TestLRUEvictionLeaksRollbackHistoryForSameIdentity(t *testing.T) {
27 - set, err := NewRelaySet(nil)
28 - if err != nil {
29 - t.Fatalf("NewRelaySet() error = %v", err)
30 - }
31 - signing := mustSigningIdentity(t)
32 - now := time.Now().UTC().Truncate(time.Microsecond)
33 - relayURL := "https://relay-replay.example"
34 -
35 - // Step 1: legitimate victim announces the NEWER descriptor. This populates
36 - // s.keyIndex[address] = t1.
37 - t1 := now
38 - t0 := now.Add(-2 * time.Minute) // strictly older but ExpiresAt is still in the future
39 - newer := mustSignedDescriptor(t, signing, "relay-replay", relayURL, t1)
40 - if accepted, _, err := set.InsertAnnounced(newer, now); err != nil || !accepted {
41 - t.Fatalf("seed announce: accepted=%v err=%v", accepted, err)
42 - }
43 -
44 - // Step 2: LRU pressure evicts the victim's slot. In production this
45 - // happens when MaxAnnouncedRelays is exceeded and the victim is the
46 - // oldest non-pinned candidate. We drive it through the same code path
47 - // enforceCapLocked uses so the assertion is faithful.
48 - set.mu.Lock()
49 - set.deleteRelayLocked(relayURL)
50 - set.mu.Unlock()
51 -
52 - // Step 3: attacker replays a strictly OLDER captured descriptor for the
53 - // same signing identity. Rollback defense MUST still reject it — the
54 - // rollback invariant is about the signing identity, not the URL slot.
55 - older := mustSignedDescriptor(t, signing, "relay-replay", relayURL, t0)
56 - accepted, _, err := set.InsertAnnounced(older, now)
57 - if accepted || err == nil {
58 - t.Fatalf(
59 - "rollback history was lost after LRU eviction: "+
60 - "InsertAnnounced(older) accepted=%v err=%v; "+
61 - "keyIndex must persist rollback history across URL-slot eviction",
62 - accepted, err,
63 - )
64 - }
65 -}
66 -
67 -// TestEnforceCapLockedSilentOverflowWhenEveryEntryPinned exercises the cap
68 -// invariant under adversarial pinning. MaxAnnouncedRelays is documented at
69 -// relaystate.go as a "hard ceiling" on the number of relay entries. The
70 -// eviction implementation only considers non-Bootstrap, non-Confirmed
71 -// entries as candidates — if every entry in the set is pinned, the candidate
72 -// list is empty, the eviction loop is a no-op, and the map silently grows
73 -// past the documented ceiling.
74 -//
75 -// This can be reached in practice when an operator bootstrap list grows past
76 -// MaxAnnouncedRelays, or when a listener confirms more relays than the cap
77 -// over a long-running session. Either way the invariant is broken and the
78 -// memory ceiling is not honored.
79 -func TestEnforceCapLockedSilentOverflowWhenEveryEntryPinned(t *testing.T) {
80 - set, err := NewRelaySet(nil)
81 - if err != nil {
82 - t.Fatalf("NewRelaySet() error = %v", err)
83 - }
84 - signing := mustSigningIdentity(t)
85 - base := time.Now().UTC().Truncate(time.Microsecond)
86 -
87 - // Populate MaxAnnouncedRelays + overflow entries, all marked Confirmed.
88 - // A correct eviction policy MUST still maintain the documented hard
89 - // ceiling; a silently-overflowing policy violates it.
90 - const overflow = 3
91 - set.mu.Lock()
92 - for i := range MaxAnnouncedRelays + overflow {
93 - url := "https://pinned-" + sprintInt(i) + ".example"
94 - set.relays[url] = RelayState{
95 - Descriptor: types.RelayDescriptor{
96 - Identity: types.Identity{Address: signing.Address},
97 - APIHTTPSAddr: url,
98 - IssuedAt: base.Add(time.Duration(i) * time.Second),
99 - ExpiresAt: base.Add(time.Hour),
100 - },
101 - Confirmed: true,
102 - LastSeenAt: base.Add(time.Duration(i) * time.Second),
103 - }
104 - }
105 - pre := len(set.relays)
106 - set.enforceCapLocked()
107 - post := len(set.relays)
108 - set.mu.Unlock()
109 -
110 - if pre <= MaxAnnouncedRelays {
111 - t.Fatalf("test setup invalid: pre=%d cap=%d", pre, MaxAnnouncedRelays)
112 - }
113 - if post > MaxAnnouncedRelays {
114 - t.Fatalf(
115 - "enforceCapLocked silently overflowed the hard ceiling: "+
116 - "post=%d cap=%d — when every entry is pinned, the eviction "+
117 - "loop has zero candidates and the set is left above "+
118 - "MaxAnnouncedRelays, violating the documented invariant",
119 - post, MaxAnnouncedRelays,
120 - )
121 - }
122 -}
portal/discovery/announce_test.go
+44 -169
@@ -4,13 +4,40 @@ import (
4 "testing"
5 "time"
6
7 + "github.com/gosuda/portal-tunnel/v2/portal/auth"
8 "github.com/gosuda/portal-tunnel/v2/types"
9 "github.com/gosuda/portal-tunnel/v2/utils"
10 )
11
12 +func mustSigningIdentity(t *testing.T) types.Identity {
13 + t.Helper()
14 + identity, err := utils.ResolveSecp256k1Identity("")
15 + if err != nil {
16 + t.Fatalf("ResolveSecp256k1Identity() error = %v", err)
17 + }
18 + return identity
19 +}
20 +
21 +func mustUnsignedDescriptor(t *testing.T, signing types.Identity, relayName, relayURL string) types.RelayDescriptor {
22 + t.Helper()
23 + now := time.Now().UTC().Truncate(time.Microsecond)
24 + return types.RelayDescriptor{
25 + Identity: types.Identity{
26 + Name: relayName,
27 + Address: signing.Address,
28 + },
29 + RelayID: relayURL,
30 + Version: 1,
31 + IssuedAt: now,
32 + ExpiresAt: now.Add(time.Hour),
33 + APIHTTPSAddr: relayURL,
34 + Discovery: true,
35 + }
36 +}
37 +
38 func mustSignedDescriptor(t *testing.T, signing types.Identity, relayName, relayURL string, issuedAt time.Time) types.RelayDescriptor {
39 t.Helper()
13 - desc, err := utils.NormalizeDescriptor(types.RelayDescriptor{
40 + signed, err := auth.SignRelayDescriptor(types.RelayDescriptor{
41 Identity: types.Identity{
42 Name: relayName,
43 Address: signing.Address,
@@ -21,114 +48,66 @@ func mustSignedDescriptor(t *testing.T, signing types.Identity, relayName, relay
48 ExpiresAt: issuedAt.Add(DiscoveryDescriptorTTL),
49 APIHTTPSAddr: relayURL,
50 Discovery: true,
24 - })
51 + }, signing.PrivateKey)
52 if err != nil {
26 - t.Fatalf("NormalizeDescriptor() error = %v", err)
27 - }
28 - signed, err := SignDescriptor(desc, signing.PrivateKey)
29 - if err != nil {
30 - t.Fatalf("SignDescriptor() error = %v", err)
53 + t.Fatalf("SignRelayDescriptor() error = %v", err)
54 }
55 return signed
56 }
57
58 func TestInsertAnnouncedAcceptsValidDescriptor(t *testing.T) {
36 - set, err := NewRelaySet(nil)
37 - if err != nil {
38 - t.Fatalf("NewRelaySet() error = %v", err)
39 - }
59 + set := NewRelaySet(nil)
60 signing := mustSigningIdentity(t)
61 now := time.Now().UTC().Truncate(time.Microsecond)
62 desc := mustSignedDescriptor(t, signing, "relay-ann", "https://relay-ann.example", now)
43 - accepted, changed, err := set.InsertAnnounced(desc, now)
44 - if err != nil {
63 + if err := set.InsertAnnounced(desc, now); err != nil {
64 t.Fatalf("InsertAnnounced() error = %v", err)
65 }
47 - if !accepted || !changed {
48 - t.Fatalf("expected accept+change, got accepted=%v changed=%v", accepted, changed)
49 - }
66 if got := set.AggregateRelays(); len(got) != 1 {
67 t.Fatalf("len(AggregateRelays()) = %d, want 1", len(got))
68 }
69 }
70
71 func TestInsertAnnouncedRejectsUnsigned(t *testing.T) {
56 - set, err := NewRelaySet(nil)
57 - if err != nil {
58 - t.Fatalf("NewRelaySet() error = %v", err)
59 - }
60 - signing := mustSigningIdentity(t)
61 - now := time.Now().UTC().Truncate(time.Microsecond)
62 - desc := mustNormalizedDescriptor(t, signing, "relay-unsigned", "https://relay-unsigned.example")
63 - if accepted, _, err := set.InsertAnnounced(desc, now); accepted || err == nil {
64 - t.Fatalf("expected unsigned reject, got accepted=%v err=%v", accepted, err)
65 - }
66 -}
67 -
68 -func TestInsertAnnouncedRejectsExpired(t *testing.T) {
69 - set, err := NewRelaySet(nil)
70 - if err != nil {
71 - t.Fatalf("NewRelaySet() error = %v", err)
72 - }
73 - signing := mustSigningIdentity(t)
74 - now := time.Now().UTC().Truncate(time.Microsecond)
75 - stale := now.Add(-2 * DiscoveryDescriptorTTL)
76 - desc := mustSignedDescriptor(t, signing, "relay-stale", "https://relay-stale.example", stale)
77 - if accepted, _, err := set.InsertAnnounced(desc, now); accepted || err == nil {
78 - t.Fatalf("expected expired reject, got accepted=%v err=%v", accepted, err)
79 - }
80 -}
81 -
82 -func TestInsertAnnouncedRejectsFutureClockSkew(t *testing.T) {
83 - set, err := NewRelaySet(nil)
84 - if err != nil {
85 - t.Fatalf("NewRelaySet() error = %v", err)
86 - }
72 + set := NewRelaySet(nil)
73 signing := mustSigningIdentity(t)
74 now := time.Now().UTC().Truncate(time.Microsecond)
89 - future := now.Add(2 * AnnounceClockSkewTolerance)
90 - desc := mustSignedDescriptor(t, signing, "relay-future", "https://relay-future.example", future)
91 - if accepted, _, err := set.InsertAnnounced(desc, now); accepted || err == nil {
92 - t.Fatalf("expected future-skew reject, got accepted=%v err=%v", accepted, err)
75 + desc := mustUnsignedDescriptor(t, signing, "relay-unsigned", "https://relay-unsigned.example")
76 + if err := set.InsertAnnounced(desc, now); err == nil {
77 + t.Fatal("expected unsigned reject")
78 }
79 }
80
81 func TestInsertAnnouncedRejectsRollback(t *testing.T) {
97 - set, err := NewRelaySet(nil)
98 - if err != nil {
99 - t.Fatalf("NewRelaySet() error = %v", err)
100 - }
82 + set := NewRelaySet(nil)
83 signing := mustSigningIdentity(t)
84 now := time.Now().UTC().Truncate(time.Microsecond)
85 relayURL := "https://relay-roll.example"
86 newer := mustSignedDescriptor(t, signing, "relay-roll", relayURL, now)
105 - if _, _, err := set.InsertAnnounced(newer, now); err != nil {
87 + if err := set.InsertAnnounced(newer, now); err != nil {
88 t.Fatalf("seed insert error = %v", err)
89 }
90 older := mustSignedDescriptor(t, signing, "relay-roll", relayURL, now.Add(-time.Minute))
109 - if accepted, _, err := set.InsertAnnounced(older, now); accepted || err == nil {
110 - t.Fatalf("expected rollback reject, got accepted=%v err=%v", accepted, err)
91 + if err := set.InsertAnnounced(older, now); err == nil {
92 + t.Fatal("expected rollback reject")
93 }
94 }
95
96 func TestInsertAnnouncedBlocksCrossIdentityTakeover(t *testing.T) {
115 - set, err := NewRelaySet(nil)
116 - if err != nil {
117 - t.Fatalf("NewRelaySet() error = %v", err)
118 - }
97 + set := NewRelaySet(nil)
98 owner := mustSigningIdentity(t)
99 attacker := mustSigningIdentity(t)
100 now := time.Now().UTC().Truncate(time.Microsecond)
101 relayURL := "https://relay-takeover.example"
102
103 ownerDesc := mustSignedDescriptor(t, owner, "relay-takeover", relayURL, now)
125 - if _, _, err := set.InsertAnnounced(ownerDesc, now); err != nil {
104 + if err := set.InsertAnnounced(ownerDesc, now); err != nil {
105 t.Fatalf("owner insert error = %v", err)
106 }
107
108 attackerDesc := mustSignedDescriptor(t, attacker, "relay-takeover", relayURL, now.Add(time.Second))
130 - if accepted, _, err := set.InsertAnnounced(attackerDesc, now); accepted || err == nil {
131 - t.Fatalf("expected takeover reject, got accepted=%v err=%v", accepted, err)
109 + if err := set.InsertAnnounced(attackerDesc, now); err == nil {
110 + t.Fatal("expected takeover reject")
111 }
112
113 states := set.AggregateRelays()
@@ -140,94 +119,6 @@ func TestInsertAnnouncedBlocksCrossIdentityTakeover(t *testing.T) {
119 }
120 }
121
143 -func TestEnforceCapLockedEvictsOldestNonPinned(t *testing.T) {
144 - set, err := NewRelaySet([]string{"https://bootstrap.example"})
145 - if err != nil {
146 - t.Fatalf("NewRelaySet() error = %v", err)
147 - }
148 - signing := mustSigningIdentity(t)
149 - base := time.Now().UTC().Truncate(time.Microsecond)
150 -
151 - // Inject MaxAnnouncedRelays + extra candidates so eviction must run.
152 - // LastSeenAt is a strict ramp so we can assert exactly which slots
153 - // the oldest-first policy removed.
154 - const extra = 5
155 - set.mu.Lock()
156 - for i := range MaxAnnouncedRelays + extra {
157 - url := "https://stub-" + sprintInt(i) + ".example"
158 - set.relays[url] = RelayState{
159 - Descriptor: types.RelayDescriptor{
160 - Identity: types.Identity{Address: signing.Address},
161 - APIHTTPSAddr: url,
162 - IssuedAt: base.Add(time.Duration(i) * time.Second),
163 - ExpiresAt: base.Add(time.Hour),
164 - },
165 - LastSeenAt: base.Add(time.Duration(i) * time.Second),
166 - }
167 - }
168 - // Pin one of the candidates as Confirmed at index 1 so we can assert
169 - // that pinned entries are NOT evicted even when they are old.
170 - pinnedURL := "https://stub-1.example"
171 - pinned := set.relays[pinnedURL]
172 - pinned.Confirmed = true
173 - set.relays[pinnedURL] = pinned
174 -
175 - preCount := len(set.relays)
176 - set.enforceCapLocked()
177 - postCount := len(set.relays)
178 -
179 - _, bootstrapPresent := set.relays["https://bootstrap.example"]
180 - _, pinnedPresent := set.relays[pinnedURL]
181 - // The oldest non-pinned candidates (index 0, then 2, 3, 4 — index 1 is
182 - // pinned) should have been evicted first. extra+1 entries are removed
183 - // because the bootstrap slot pushes total over the cap by one.
184 - _, oldest0Present := set.relays["https://stub-0.example"]
185 - _, oldest2Present := set.relays["https://stub-2.example"]
186 - _, newestPresent := set.relays["https://stub-"+sprintInt(MaxAnnouncedRelays+extra-1)+".example"]
187 - set.mu.Unlock()
188 -
189 - if preCount <= MaxAnnouncedRelays {
190 - t.Fatalf("test setup invalid: preCount=%d cap=%d", preCount, MaxAnnouncedRelays)
191 - }
192 - if postCount > MaxAnnouncedRelays {
193 - t.Fatalf("postCount=%d exceeds cap=%d", postCount, MaxAnnouncedRelays)
194 - }
195 - if !bootstrapPresent {
196 - t.Fatal("bootstrap entry must survive LRU eviction")
197 - }
198 - if !pinnedPresent {
199 - t.Fatal("Confirmed entry must survive LRU eviction even if old")
200 - }
201 - if oldest0Present {
202 - t.Fatal("oldest non-pinned entry (stub-0) must be evicted first")
203 - }
204 - if oldest2Present {
205 - t.Fatal("second-oldest non-pinned entry (stub-2) must also be evicted")
206 - }
207 - if !newestPresent {
208 - t.Fatal("newest entry must survive LRU eviction")
209 - }
210 -}
211 -
212 -func sprintInt(n int) string {
213 - if n == 0 {
214 - return "0"
215 - }
216 - digits := make([]byte, 0, 6)
217 - negative := n < 0
218 - if negative {
219 - n = -n
220 - }
221 - for n > 0 {
222 - digits = append([]byte{byte('0' + n%10)}, digits...)
223 - n /= 10
224 - }
225 - if negative {
226 - return "-" + string(digits)
227 - }
228 - return string(digits)
229 -}
230 -
122 func TestAnnounceLimiterAllowsBurstThenThrottles(t *testing.T) {
123 limiter := NewAnnounceLimiter(60, 5) // 1/sec sustained, burst 5
124 for i := range 5 {
@@ -242,19 +133,3 @@ func TestAnnounceLimiterAllowsBurstThenThrottles(t *testing.T) {
133 t.Fatal("different IP should have its own bucket")
134 }
135 }
245 -
246 -func TestAnnounceLimiterRefillsOverTime(t *testing.T) {
247 - limiter := NewAnnounceLimiter(60, 1) // 1/sec sustained, burst 1
248 - clock := time.Now()
249 - limiter.clock = func() time.Time { return clock }
250 - if !limiter.Allow("10.0.0.1") {
251 - t.Fatal("first request should be allowed")
252 - }
253 - if limiter.Allow("10.0.0.1") {
254 - t.Fatal("second immediate request should be throttled")
255 - }
256 - clock = clock.Add(2 * time.Second)
257 - if !limiter.Allow("10.0.0.1") {
258 - t.Fatal("after refill the request should be allowed again")
259 - }
260 -}
portal/discovery/policy.go
+3 -3
@@ -30,7 +30,7 @@ func (p DefaultRelayPolicy) SelectAggregate(states []RelayState) []RelayState {
30 if state.Banned {
31 continue
32 }
33 - if !state.Bootstrap && !state.hasDescriptor() {
33 + if !state.Bootstrap && !state.hasObservedDescriptor() {
34 continue
35 }
36 out = append(out, state)
@@ -65,10 +65,10 @@ func (p DefaultRelayPolicy) SelectPriority(states []RelayState, clientState Clie
65 explicit := make([]string, 0, len(clientState.ExplicitRelayURLs))
66 autoPool := make([]RelayState, 0, len(selected))
67 for _, state := range selected {
68 - if clientState.RequireUDP && state.hasDescriptor() && !state.Descriptor.SupportsUDP {
68 + if clientState.RequireUDP && state.hasObservedDescriptor() && !state.Descriptor.SupportsUDP {
69 continue
70 }
71 - if clientState.RequireTCP && state.hasDescriptor() && !state.Descriptor.SupportsTCP {
71 + if clientState.RequireTCP && state.hasObservedDescriptor() && !state.Descriptor.SupportsTCP {
72 continue
73 }
74 relayURL := state.Descriptor.APIHTTPSAddr
portal/discovery/policy_test.go
+5 -32
@@ -5,6 +5,7 @@ import (
5 "testing"
6 "time"
7
8 + "github.com/gosuda/portal-tunnel/v2/portal/auth"
9 "github.com/gosuda/portal-tunnel/v2/types"
10 "github.com/gosuda/portal-tunnel/v2/utils"
11 )
@@ -17,7 +18,7 @@ func mustPolicyRelayDescriptor(t *testing.T, relayName, relayURL string) types.R
18 t.Fatalf("ResolveSecp256k1Identity() error = %v", err)
19 }
20 now := time.Now().UTC()
20 - desc, err := utils.NormalizeDescriptor(types.RelayDescriptor{
21 + signed, err := auth.SignRelayDescriptor(types.RelayDescriptor{
22 Identity: types.Identity{
23 Name: relayName,
24 Address: signing.Address,
@@ -28,13 +29,9 @@ func mustPolicyRelayDescriptor(t *testing.T, relayName, relayURL string) types.R
29 ExpiresAt: now.Add(time.Hour),
30 APIHTTPSAddr: relayURL,
31 Discovery: true,
31 - })
32 - if err != nil {
33 - t.Fatalf("NormalizeDescriptor() error = %v", err)
34 - }
35 - signed, err := SignDescriptor(desc, signing.PrivateKey)
32 + }, signing.PrivateKey)
33 if err != nil {
37 - t.Fatalf("SignDescriptor() error = %v", err)
34 + t.Fatalf("SignRelayDescriptor() error = %v", err)
35 }
36 return signed
37 }
@@ -116,7 +113,7 @@ func TestSelectAggregateKeepsCollectedRelayEvenWhenNotAdvertisable(t *testing.T)
113 policy := DefaultRelayPolicy{}
114 state := RelayState{
115 Descriptor: mustPolicyRelayDescriptor(t, "relay-a", "https://relay-a.example"),
119 - LastSeenAt: time.Now().UTC().Add(-DiscoveryHintRetentionTTL).Add(-time.Hour),
116 + LastSeenAt: time.Now().UTC().Add(-31 * 24 * time.Hour),
117 }
118 state.Descriptor.Discovery = false
119 state.Descriptor.ExpiresAt = time.Now().UTC().Add(-time.Second)
@@ -157,30 +154,6 @@ func TestSelectAggregateSkipsBannedBootstrapRelay(t *testing.T) {
154 }
155 }
156
160 -func TestSelectConfirmedKeepsOnlyConfirmedAggregateRelays(t *testing.T) {
161 - policy := DefaultRelayPolicy{}
162 - confirmed := confirmedPolicyRelayState(t, "relay-confirmed", "https://relay-confirmed.example")
163 - hinted := RelayState{
164 - Descriptor: mustPolicyRelayDescriptor(t, "relay-hinted", "https://relay-hinted.example"),
165 - LastSeenAt: time.Now().UTC(),
166 - }
167 - bannedConfirmed := confirmedPolicyRelayState(t, "relay-banned", "https://relay-banned.example")
168 - bannedConfirmed.Banned = true
169 -
170 - selected := policy.SelectConfirmed([]RelayState{
171 - hinted,
172 - confirmed,
173 - bannedConfirmed,
174 - })
175 -
176 - if len(selected) != 1 {
177 - t.Fatalf("len(selected) = %d, want 1", len(selected))
178 - }
179 - if got := selected[0].Descriptor.APIHTTPSAddr; got != confirmed.Descriptor.APIHTTPSAddr {
180 - t.Fatalf("selected[0] = %q, want confirmed relay %q", got, confirmed.Descriptor.APIHTTPSAddr)
181 - }
182 -}
183 -
157 func TestSelectPriorityColdStartSelectsEligibleRelay(t *testing.T) {
158 policy := DefaultRelayPolicy{}
159 relayA := "https://relay-a.example"
portal/discovery/refresher.go
+48 -62
@@ -3,9 +3,6 @@ package discovery
3 import (
4 "context"
5 "crypto/tls"
6 - "crypto/x509"
7 - "errors"
8 - "net"
6 "net/http"
7 "net/url"
8 "time"
@@ -31,40 +28,16 @@ type Refresher struct {
28 relaySet *RelaySet
29 httpClient *http.Client
30 overlay OverlayRuntime
34 - sourceBaseURL *url.URL
31 directRecoveryFailures int
32 }
33
38 -func NewRefresher(relaySet *RelaySet, rootCAPEM []byte, overlay OverlayRuntime, sourceBaseURL string) (*Refresher, error) {
39 - if relaySet == nil {
40 - return nil, errors.New("relay set is required")
41 - }
42 - var rootCAs *x509.CertPool
43 - if len(rootCAPEM) > 0 {
44 - rootCAs = x509.NewCertPool()
45 - if !rootCAs.AppendCertsFromPEM(rootCAPEM) {
46 - return nil, errors.New("failed to parse relay root ca")
47 - }
48 - }
49 - var parsedSourceBaseURL *url.URL
50 - if sourceBaseURL != "" {
51 - normalizedSourceBaseURL, err := utils.NormalizeRelayURL(sourceBaseURL)
52 - if err != nil {
53 - return nil, err
54 - }
55 - parsed, err := url.Parse(normalizedSourceBaseURL)
56 - if err != nil {
57 - return nil, err
58 - }
59 - parsedSourceBaseURL = parsed
60 - }
34 +func NewRefresher(relaySet *RelaySet, overlay OverlayRuntime) *Refresher {
35 return &Refresher{
36 relaySet: relaySet,
37 httpClient: &http.Client{
38 Transport: &http.Transport{
39 TLSClientConfig: &tls.Config{
40 MinVersion: tls.VersionTLS12,
67 - RootCAs: rootCAs,
41 NextProtos: []string{"http/1.1"},
42 },
43 ForceAttemptHTTP2: false,
@@ -72,12 +45,11 @@ func NewRefresher(relaySet *RelaySet, rootCAPEM []byte, overlay OverlayRuntime,
45 Timeout: defaultRequestTimeout,
46 },
47 overlay: overlay,
75 - sourceBaseURL: parsedSourceBaseURL,
48 directRecoveryFailures: defaultRecoveryFailures,
77 - }, nil
49 + }
50 }
51
80 -func (r *Refresher) Refresh(ctx context.Context, extraSourceHosts ...string) error {
52 +func (r *Refresher) Refresh(ctx context.Context, self types.RelayDescriptor) error {
53 if r.overlay != nil {
54 if err := r.refreshOverlay(ctx); err != nil && ctx.Err() == nil {
55 log.Warn().
@@ -88,12 +60,52 @@ func (r *Refresher) Refresh(ctx context.Context, extraSourceHosts ...string) err
60 return ctx.Err()
61 }
62 }
91 - return r.refreshHTTPS(ctx, extraSourceHosts)
63 + if err := r.refreshHTTPS(ctx); err != nil {
64 + return err
65 + }
66 + return r.announceSelf(ctx, self)
67 +}
68 +
69 +func (r *Refresher) announceSelf(ctx context.Context, descriptor types.RelayDescriptor) error {
70 + req := types.DiscoveryAnnounceRequest{
71 + ProtocolVersion: types.DiscoveryVersion,
72 + Descriptor: descriptor,
73 + }
74 + for _, relayURL := range r.relaySet.BootstrapRelayURLs() {
75 + if relayURL == descriptor.APIHTTPSAddr {
76 + continue
77 + }
78 + baseURL, err := url.Parse(relayURL)
79 + if err != nil {
80 + log.Warn().
81 + Err(err).
82 + Str("relay", relayURL).
83 + Msg("relay discovery announce target skipped")
84 + continue
85 + }
86 + if utils.IsLocalRelayHost(baseURL.Hostname()) {
87 + continue
88 + }
89 +
90 + if err := utils.HTTPDoAPIPath(ctx, r.httpClient, baseURL, http.MethodPost, types.PathDiscoveryAnnounce, req, nil, nil); err != nil {
91 + if ctx.Err() != nil {
92 + return ctx.Err()
93 + }
94 + log.Warn().
95 + Err(err).
96 + Str("relay", relayURL).
97 + Msg("relay discovery announce failed")
98 + }
99 + }
100 + return nil
101 }
102
94 -func (r *Refresher) refreshHTTPS(ctx context.Context, extraSourceHosts []string) error {
103 +func (r *Refresher) refreshHTTPS(ctx context.Context) error {
104 r.relaySet.mu.RLock()
96 - states := r.relaySet.relayStatesLocked()
105 + states := make([]RelayState, 0, len(r.relaySet.relays))
106 + for _, state := range r.relaySet.relays {
107 + states = append(states, state)
108 + }
109 r.relaySet.mu.RUnlock()
110
111 now := time.Now().UTC()
@@ -101,7 +113,7 @@ func (r *Refresher) refreshHTTPS(ctx context.Context, extraSourceHosts []string)
113 if state.Banned {
114 continue
115 }
104 - if !state.hasDescriptor() {
116 + if !state.hasObservedDescriptor() {
117 if !state.Bootstrap {
118 continue
119 }
@@ -110,7 +122,7 @@ func (r *Refresher) refreshHTTPS(ctx context.Context, extraSourceHosts []string)
122 continue
123 }
124 }
113 - if state.hasDescriptor() && !state.Descriptor.Discovery {
125 + if state.hasObservedDescriptor() && !state.Descriptor.Discovery {
126 continue
127 }
128
@@ -159,32 +171,6 @@ func (r *Refresher) refreshHTTPS(ctx context.Context, extraSourceHosts []string)
171 }
172 r.relaySet.RecordDiscoveryRTT(relayURL, time.Since(startedAt), measuredAt)
173 }
162 -
163 - for _, sourceHost := range extraSourceHosts {
164 - if r.sourceBaseURL == nil {
165 - break
166 - }
167 - sourceHost = utils.NormalizeHostname(sourceHost)
168 - if sourceHost == "" {
169 - continue
170 - }
171 - baseURL := *r.sourceBaseURL
172 - baseURL.Host = sourceHost
173 - if port := r.sourceBaseURL.Port(); port != "" {
174 - baseURL.Host = net.JoinHostPort(sourceHost, port)
175 - }
176 -
177 - var resp types.DiscoveryResponse
178 - if err := utils.HTTPDoAPIPath(ctx, r.httpClient, &baseURL, http.MethodGet, types.PathDiscovery, nil, nil, &resp); err != nil {
179 - if ctx.Err() != nil {
180 - return ctx.Err()
181 - }
182 - continue
183 - }
184 - if _, err := r.relaySet.ApplyRelayDiscoveryResponse("", resp, time.Now().UTC()); err != nil {
185 - continue
186 - }
187 - }
174 return nil
175 }
176
portal/discovery/relayset.go
+122 -160
@@ -3,15 +3,14 @@ package discovery
3 import (
4 "errors"
5 "fmt"
6 - "net/http"
6 "reflect"
7 "sort"
8 "strings"
9 "sync"
10 "time"
11
12 + "github.com/gosuda/portal-tunnel/v2/portal/auth"
13 "github.com/gosuda/portal-tunnel/v2/types"
14 - "github.com/gosuda/portal-tunnel/v2/utils"
14 )
15
16 // RelaySet owns the shared relay discovery view: configured bootstrap relay URLs,
@@ -33,12 +32,12 @@ import (
32 // last URL slot for an identity (via LRU or explicit removal) MUST NOT forget
33 // the rollback anchor, otherwise a captured older-but-unexpired descriptor
34 // could be replayed after eviction. Tombstones expire once the replay window
36 -// closes, i.e. once now > IssuedAt + AnnounceMaxValidity — by that time any
37 -// descriptor whose IssuedAt is ≤ the tombstoned value is strictly expired and
35 +// closes, i.e. once now > IssuedAt + AnnounceMaxValidity. By that time any
36 +// descriptor whose IssuedAt is at or before the tombstoned value is expired and
37 // cannot pass the announce validity check regardless.
38 //
39 // Both maps must always be read and written under s.mu. Mutators come in two
41 -// flavors: public methods that own the lock end-to-end, and *Locked helpers
40 +// flavors: public methods that own the lock end-to-end, and *Locked methods
41 // that assume the caller already holds s.mu as a write lock and never re-
42 // acquire it themselves. This convention prevents nested-locking deadlocks
43 // (notably from ApplyRelayDiscoveryResponse, which holds the write lock for
@@ -53,30 +52,21 @@ type RelaySet struct {
52 // keyIndexEntry records the rollback anchor for a signing identity.
53 // IssuedAt is the newest descriptor IssuedAt the set has ever accepted
54 // for this identity. TombstoneUntil is the wall-clock time at which the
56 -// rollback anchor may safely be forgotten — after that point, any
55 +// rollback anchor may safely be forgotten. After that point, any
56 // replayable descriptor with an older IssuedAt is itself expired.
57 type keyIndexEntry struct {
58 IssuedAt time.Time
59 TombstoneUntil time.Time
60 }
61
63 -func NewRelaySet(bootstrapRelayURLs []string) (*RelaySet, error) {
62 +func NewRelaySet(bootstrapRelayURLs []string) *RelaySet {
63 set := &RelaySet{
64 relays: make(map[string]RelayState),
65 keyIndex: make(map[string]keyIndexEntry),
66 policy: DefaultRelayPolicy{},
67 }
69 - if err := set.SetBootstrapRelayURLs(bootstrapRelayURLs); err != nil {
70 - return nil, err
71 - }
72 - return set, nil
73 -}
74 -
75 -// keyIndexAddress returns the lower-cased EVM address used as the keyIndex
76 -// key for a given relay state. Empty for stub entries that carry no signed
77 -// descriptor (e.g. bootstrap URL placeholders before the first refresh).
78 -func keyIndexAddress(state RelayState) string {
79 - return strings.ToLower(strings.TrimSpace(state.Descriptor.Address))
68 + set.SetBootstrapRelayURLs(bootstrapRelayURLs)
69 + return set
70 }
71
72 // upsertDescriptorLocked applies a fully-merged RelayState to s.relays and
@@ -105,7 +95,7 @@ func (s *RelaySet) upsertDescriptorLocked(record RelayState, now time.Time, allo
95 if relayURL == "" {
96 return false
97 }
108 - address := keyIndexAddress(record)
98 + address := strings.ToLower(strings.TrimSpace(record.Descriptor.Address))
99 if address != "" {
100 if prev, ok := s.keyIndex[address]; ok {
101 // Stale tombstone: no replayable descriptor could still be
@@ -120,7 +110,7 @@ func (s *RelaySet) upsertDescriptorLocked(record RelayState, now time.Time, allo
110 }
111 if !allowCrossIdentityTakeover {
112 if existing, ok := s.relays[relayURL]; ok {
123 - existingAddress := keyIndexAddress(existing)
113 + existingAddress := strings.ToLower(strings.TrimSpace(existing.Descriptor.Address))
114 if existingAddress != "" && address != "" && existingAddress != address {
115 if !existing.Descriptor.ExpiresAt.IsZero() && existing.Descriptor.ExpiresAt.After(now) {
116 return false
@@ -148,36 +138,6 @@ func (s *RelaySet) upsertDescriptorLocked(record RelayState, now time.Time, allo
138 return true
139 }
140
151 -// deleteRelayLocked removes a URL slot from s.relays. The keyIndex tombstone
152 -// is intentionally NOT dropped here: the rollback anchor must outlive the
153 -// URL slot so that LRU eviction cannot be used as a laundering step for a
154 -// captured older-but-unexpired descriptor from the same signing identity.
155 -// Stale tombstones are swept by pruneKeyIndexLocked, called from
156 -// enforceCapLocked after every insert. The caller MUST already hold s.mu
157 -// as a write lock.
158 -func (s *RelaySet) deleteRelayLocked(relayURL string) {
159 - if _, ok := s.relays[relayURL]; !ok {
160 - return
161 - }
162 - delete(s.relays, relayURL)
163 -}
164 -
165 -// pruneKeyIndexLocked drops keyIndex tombstones whose replay-window has
166 -// closed. A tombstone at `now.After(entry.TombstoneUntil)` cannot gate any
167 -// live descriptor: the oldest replayable descriptor from the same identity
168 -// would itself be expired (since honest announces cap validity at
169 -// AnnounceMaxValidity). Callers MUST already hold s.mu as a write lock.
170 -func (s *RelaySet) pruneKeyIndexLocked(now time.Time) {
171 - for address, entry := range s.keyIndex {
172 - if entry.TombstoneUntil.IsZero() {
173 - continue
174 - }
175 - if now.After(entry.TombstoneUntil) {
176 - delete(s.keyIndex, address)
177 - }
178 - }
179 -}
180 -
141 func (s *RelaySet) SetRelayPolicy(policy RelayPolicy) {
142 if policy == nil {
143 policy = DefaultRelayPolicy{}
@@ -187,7 +147,7 @@ func (s *RelaySet) SetRelayPolicy(policy RelayPolicy) {
147 s.policy = policy
148 }
149
190 -func (s *RelaySet) SetBootstrapRelayURLs(inputs []string) error {
150 +func (s *RelaySet) SetBootstrapRelayURLs(inputs []string) {
151 s.mu.Lock()
152 defer s.mu.Unlock()
153
@@ -199,8 +159,8 @@ func (s *RelaySet) SetBootstrapRelayURLs(inputs []string) error {
159 for key, state := range s.relays {
160 _, bootstrap := keep[key]
161 state.Bootstrap = bootstrap
202 - if !state.Bootstrap && !state.hasDescriptor() && !state.Banned && state.consecutiveFailures == 0 {
203 - s.deleteRelayLocked(key)
162 + if !state.Bootstrap && !state.hasObservedDescriptor() && !state.Banned && state.consecutiveFailures == 0 {
163 + delete(s.relays, key)
164 continue
165 }
166
@@ -212,43 +172,54 @@ func (s *RelaySet) SetBootstrapRelayURLs(inputs []string) error {
172 continue
173 }
174
215 - state := newRelayStateFromURL(relayURL)
175 + state := newRelayState(relayURL)
176 state.Bootstrap = true
177 s.relays[relayURL] = state
178 }
219 - return nil
179 }
180
181 func (s *RelaySet) AggregateRelays() []RelayState {
182 s.mu.RLock()
224 - defer s.mu.RUnlock()
183 + states := make([]RelayState, 0, len(s.relays))
184 + for _, state := range s.relays {
185 + states = append(states, state)
186 + }
187 + policy := s.policy
188 + s.mu.RUnlock()
189
226 - return s.policy.SelectAggregate(s.relayStatesLocked())
190 + return policy.SelectAggregate(states)
191 }
192
193 func (s *RelaySet) ConfirmedRelays() []RelayState {
194 s.mu.RLock()
231 - defer s.mu.RUnlock()
195 + states := make([]RelayState, 0, len(s.relays))
196 + for _, state := range s.relays {
197 + states = append(states, state)
198 + }
199 + policy := s.policy
200 + s.mu.RUnlock()
201
233 - return s.policy.SelectConfirmed(s.relayStatesLocked())
202 + return policy.SelectConfirmed(states)
203 }
204
205 func (s *RelaySet) PriorityRelays(clientState ClientState) []string {
206 s.mu.RLock()
238 - defer s.mu.RUnlock()
207 + states := make([]RelayState, 0, len(s.relays))
208 + for _, state := range s.relays {
209 + states = append(states, state)
210 + }
211 + policy := s.policy
212 + s.mu.RUnlock()
213
240 - return s.policy.SelectPriority(s.relayStatesLocked(), clientState)
214 + return policy.SelectPriority(states, clientState)
215 }
216
217 func (s *RelaySet) OverlayPeerStates() []RelayState {
244 - s.mu.RLock()
245 - states := s.relayStatesLocked()
246 - s.mu.RUnlock()
247 -
218 now := time.Now().UTC()
249 - out := make([]RelayState, 0, len(states))
250 - for _, state := range states {
251 - if state.Banned || !state.hasDescriptor() || !state.Descriptor.ExpiresAt.After(now) || !state.Descriptor.SupportsOverlayPeer {
219 + s.mu.RLock()
220 + out := make([]RelayState, 0, len(s.relays))
221 + for _, state := range s.relays {
222 + if state.Banned || !state.hasObservedDescriptor() || !state.Descriptor.ExpiresAt.After(now) || !state.Descriptor.SupportsOverlayPeer {
223 continue
224 }
225 if state.Descriptor.WireGuardPublicKey == "" ||
@@ -258,87 +229,71 @@ func (s *RelaySet) OverlayPeerStates() []RelayState {
229 }
230 out = append(out, state)
231 }
232 + s.mu.RUnlock()
233 if len(out) == 0 {
234 return nil
235 }
236 return out
237 }
238
267 -func (s *RelaySet) Descriptors() []types.RelayDescriptor {
239 +// BootstrapRelayURLs returns configured bootstrap discovery endpoints that
240 +// can receive this relay's periodic self-announce.
241 +func (s *RelaySet) BootstrapRelayURLs() []string {
242 s.mu.RLock()
269 - states := s.relayStatesLocked()
270 - s.mu.RUnlock()
271 -
272 - now := time.Now().UTC()
273 - out := make([]types.RelayDescriptor, 0, len(states))
274 - for _, state := range states {
275 - if state.Banned || !state.hasDescriptor() || !state.Descriptor.Discovery {
243 + out := make([]string, 0, len(s.relays))
244 + for _, state := range s.relays {
245 + if state.Banned || !state.Bootstrap {
246 continue
247 }
278 - desc := state.Descriptor
279 - if !desc.ExpiresAt.After(now) {
280 - if state.LastSeenAt.IsZero() || !state.LastSeenAt.After(now.Add(-DiscoveryHintRetentionTTL)) {
281 - continue
282 - }
283 -
284 - // Keep stale relay hints flowing through discovery so the mesh converges
285 - // on a large shared relay set. Local listener confirmation and direct
286 - // refresh retry state are tracked separately.
287 - desc.ExpiresAt = now.Add(DiscoveryDescriptorTTL)
248 + if state.hasObservedDescriptor() && !state.Descriptor.Discovery {
249 + continue
250 }
289 - out = append(out, desc)
251 + relayURL := strings.TrimSpace(state.Descriptor.APIHTTPSAddr)
252 + if relayURL == "" {
253 + continue
254 + }
255 + out = append(out, relayURL)
256 }
257 + s.mu.RUnlock()
258 if len(out) == 0 {
259 return nil
260 }
261 return out
262 }
263
297 -func (s *RelaySet) ServeDiscovery(w http.ResponseWriter, r *http.Request, local ...types.RelayDescriptor) {
298 - if !utils.RequireMethod(w, r, http.MethodGet) {
299 - return
300 - }
301 -
302 - known := s.Descriptors()
303 - relays := make([]types.RelayDescriptor, 0, len(local)+len(known))
304 - seen := make(map[string]struct{}, len(local)+len(known))
305 - add := func(descriptor types.RelayDescriptor) {
306 - relayURL := descriptor.APIHTTPSAddr
264 +func (s *RelaySet) Descriptors(self types.RelayDescriptor) []types.RelayDescriptor {
265 + now := time.Now().UTC()
266 + out := make([]types.RelayDescriptor, 0, 1)
267 + seen := make(map[string]struct{})
268 + add := func(desc types.RelayDescriptor) {
269 + relayURL := desc.APIHTTPSAddr
270 if relayURL == "" {
271 return
272 }
273 if _, ok := seen[relayURL]; ok {
274 return
275 }
276 + if !desc.ExpiresAt.After(now) {
277 + return
278 + }
279 seen[relayURL] = struct{}{}
314 - relays = append(relays, descriptor)
280 + out = append(out, desc)
281 }
282
317 - for _, descriptor := range local {
318 - add(descriptor)
283 + if self.APIHTTPSAddr != "" && self.Discovery && self.ExpiresAt.After(now) {
284 + add(self)
285 }
320 - for _, descriptor := range known {
321 - add(descriptor)
322 - }
323 -
324 - utils.WriteAPIData(w, http.StatusOK, types.DiscoveryResponse{
325 - ProtocolVersion: types.DiscoveryVersion,
326 - GeneratedAt: time.Now().UTC(),
327 - Relays: relays,
328 - })
329 -}
330 -
331 -func (s *RelaySet) relayStatesLocked() []RelayState {
332 - out := make([]RelayState, 0, len(s.relays))
286 + s.mu.RLock()
287 for _, state := range s.relays {
334 - out = append(out, state)
288 + if state.Banned || !state.hasObservedDescriptor() || !state.Descriptor.Discovery {
289 + continue
290 + }
291 + add(state.Descriptor)
292 }
293 + s.mu.RUnlock()
294 if len(out) == 0 {
295 return nil
296 }
339 - sort.Slice(out, func(i, j int) bool {
340 - return out[i].Descriptor.APIHTTPSAddr < out[j].Descriptor.APIHTTPSAddr
341 - })
297 return out
298 }
299
@@ -348,7 +303,7 @@ func (s *RelaySet) BanRelayURL(relayURL string) {
303
304 state, ok := s.relays[relayURL]
305 if !ok {
351 - state = newRelayStateFromURL(relayURL)
306 + state = newRelayState(relayURL)
307 }
308 state = s.policy.OnBanned(state)
309 s.relays[relayURL] = state
@@ -360,7 +315,7 @@ func (s *RelaySet) ConfirmRelayURL(relayURL string) {
315
316 state, ok := s.relays[relayURL]
317 if !ok {
363 - state = newRelayStateFromURL(relayURL)
318 + state = newRelayState(relayURL)
319 }
320 state = s.policy.OnConfirmed(state)
321 s.relays[relayURL] = state
@@ -396,17 +351,21 @@ func (s *RelaySet) ApplyRelayDiscoveryResponse(targetURL string, resp types.Disc
351 add := func(descriptor types.RelayDescriptor) {
352 // Cryptographic gate: every gossiped descriptor must carry a valid
353 // signature. Unsigned or invalid-signature descriptors are dropped
399 - // silently — they cannot poison the local relay set, and other peers
354 + // silently; they cannot poison the local relay set, and other peers
355 // will reach the same verdict independently. This is the sole global
356 // trust gate under unconditional propagation, so it is mandatory.
402 - if _, verifyErr := VerifyDescriptor(descriptor); verifyErr != nil {
357 + verified, verifyErr := auth.VerifyRelayDescriptor(descriptor)
358 + if verifyErr != nil {
359 return
360 }
405 - relayState, err := newRelayState(descriptor, now)
406 - if err != nil {
361 + if err := validateRelayDescriptorFreshness(verified, now); err != nil {
362 return
363 }
409 - relayURL := relayState.Descriptor.APIHTTPSAddr
364 + relayState := RelayState{
365 + Descriptor: verified,
366 + LastSeenAt: now,
367 + }
368 + relayURL := verified.APIHTTPSAddr
369 if relayURL == "" {
370 return
371 }
@@ -501,50 +460,32 @@ func (s *RelaySet) RecordDiscoveryRTT(relayURL string, rtt time.Duration, measur
460 // than AnnounceMaxValidity).
461 // 3. Local merge preserves Bootstrap, Confirmed, Banned, telemetry, and
462 // direct-refresh retry state from any pre-existing entry at the same URL.
504 -// 4. The shared upsertDescriptorLocked helper enforces the
463 +// 4. The shared upsertDescriptorLocked method enforces the
464 // monotonic-IssuedAt-per-key rollback guard and the cross-identity
506 -// URL-takeover guard. Announce never grants takeover authority — only
465 +// URL-takeover guard. Announce never grants takeover authority; only
466 // direct authoritative refresh can do that.
467 // 5. After a successful upsert, the LRU cap is enforced; bootstrap and
468 // listener-confirmed entries are pinned.
469 //
511 -// Returns (accepted, changed, err): accepted=true iff the descriptor was
512 -// stored (or was an idempotent refresh). changed=true iff s.relays was
513 -// mutated. The error categories are exported as Err* sentinels so callers
514 -// can map to HTTP statuses.
515 -func (s *RelaySet) InsertAnnounced(desc types.RelayDescriptor, now time.Time) (accepted bool, changed bool, err error) {
470 +// Returns nil iff the descriptor was stored or idempotently refreshed.
471 +func (s *RelaySet) InsertAnnounced(desc types.RelayDescriptor, now time.Time) error {
472 if now.IsZero() {
473 now = time.Now().UTC()
474 } else {
475 now = now.UTC()
476 }
477
522 - if _, verifyErr := VerifyDescriptor(desc); verifyErr != nil {
523 - return false, false, verifyErr
524 - }
525 - normalized, err := utils.NormalizeDescriptor(desc)
478 + normalized, err := auth.VerifyRelayDescriptor(desc)
479 if err != nil {
527 - return false, false, fmt.Errorf("normalize announced descriptor: %w", err)
528 - }
529 - if normalized.IssuedAt.IsZero() {
530 - return false, false, errors.New("announced descriptor missing issued_at")
531 - }
532 - if normalized.ExpiresAt.IsZero() {
533 - return false, false, errors.New("announced descriptor missing expires_at")
480 + return err
481 }
535 - if !normalized.ExpiresAt.After(now) {
536 - return false, false, errors.New("announced descriptor already expired")
537 - }
538 - if normalized.IssuedAt.After(now.Add(AnnounceClockSkewTolerance)) {
539 - return false, false, errors.New("announced descriptor is too far in the future")
540 - }
541 - if normalized.ExpiresAt.Sub(normalized.IssuedAt) > AnnounceMaxValidity {
542 - return false, false, errors.New("announced descriptor validity window exceeds maximum")
482 + if err := validateRelayDescriptorFreshness(normalized, now); err != nil {
483 + return err
484 }
485
545 - record, err := newRelayState(normalized, now)
546 - if err != nil {
547 - return false, false, err
486 + record := RelayState{
487 + Descriptor: normalized,
488 + LastSeenAt: now,
489 }
490
491 s.mu.Lock()
@@ -566,24 +507,45 @@ func (s *RelaySet) InsertAnnounced(desc types.RelayDescriptor, now time.Time) (a
507 }
508
509 if !s.upsertDescriptorLocked(record, now, false) {
569 - return false, false, errors.New("announced descriptor rejected by rollback or takeover guard")
510 + return errors.New("announced descriptor rejected by rollback or takeover guard")
511 }
512
513 s.enforceCapLocked()
573 - return true, true, nil
514 + return nil
515 +}
516 +
517 +func validateRelayDescriptorFreshness(desc types.RelayDescriptor, now time.Time) error {
518 + if desc.IssuedAt.IsZero() {
519 + return errors.New("relay descriptor missing issued_at")
520 + }
521 + if !desc.ExpiresAt.After(now) {
522 + return errors.New("relay descriptor already expired")
523 + }
524 + if desc.IssuedAt.After(now.Add(AnnounceClockSkewTolerance)) {
525 + return errors.New("relay descriptor is too far in the future")
526 + }
527 + if desc.ExpiresAt.Sub(desc.IssuedAt) > AnnounceMaxValidity {
528 + return errors.New("relay descriptor validity window exceeds maximum")
529 + }
530 + return nil
531 }
532
533 // enforceCapLocked trims s.relays back to MaxAnnouncedRelays using a
534 // two-tier eviction strategy: non-Bootstrap non-Confirmed entries are
535 // evicted first (oldest by LastSeenAt), then non-Bootstrap Confirmed
579 -// entries as a last resort. Bootstrap entries are absolutely pinned —
580 -// an operator misconfig that lists more than MaxAnnouncedRelays bootstraps
536 +// entries as a last resort. Bootstrap entries are absolutely pinned.
537 +// An operator misconfig that lists more than MaxAnnouncedRelays bootstraps
538 // is surfaced by the resulting overflow rather than silently violating
539 // operator intent. Tombstone keyIndex entries whose replay window has
540 // closed are swept opportunistically. The caller MUST already hold s.mu
541 // as a write lock.
542 func (s *RelaySet) enforceCapLocked() {
586 - s.pruneKeyIndexLocked(time.Now().UTC())
543 + now := time.Now().UTC()
544 + for address, entry := range s.keyIndex {
545 + if !entry.TombstoneUntil.IsZero() && now.After(entry.TombstoneUntil) {
546 + delete(s.keyIndex, address)
547 + }
548 + }
549 if len(s.relays) <= MaxAnnouncedRelays {
550 return
551 }
@@ -604,7 +566,7 @@ func (s *RelaySet) enforceCapLocked() {
566 })
567 }
568 sort.Slice(candidates, func(i, j int) bool {
607 - // Non-confirmed entries evict first — confirmed is the last-resort
569 + // Non-confirmed entries evict first; confirmed is the last-resort
570 // tier. Within each tier, oldest LastSeenAt evicts first.
571 if candidates[i].confirmed != candidates[j].confirmed {
572 return !candidates[i].confirmed
@@ -615,7 +577,7 @@ func (s *RelaySet) enforceCapLocked() {
577 if len(s.relays) <= MaxAnnouncedRelays {
578 return
579 }
618 - s.deleteRelayLocked(c.url)
580 + delete(s.relays, c.url)
581 }
582 }
583
portal/discovery/relayset_test.go
+11 -73
@@ -8,10 +8,7 @@ import (
8 )
9
10 func TestApplyRelayDiscoveryResponsePreservesBootstrapFlag(t *testing.T) {
11 - set, err := NewRelaySet([]string{"https://relay-a.example"})
12 - if err != nil {
13 - t.Fatalf("NewRelaySet() error = %v", err)
14 - }
11 + set := NewRelaySet([]string{"https://relay-a.example"})
12
13 desc := mustPolicyRelayDescriptor(t, "relay-a", "https://relay-a.example")
14 if _, err := set.ApplyRelayDiscoveryResponse(desc.APIHTTPSAddr, types.DiscoveryResponse{
@@ -30,11 +27,8 @@ func TestApplyRelayDiscoveryResponsePreservesBootstrapFlag(t *testing.T) {
27 }
28 }
29
33 -func TestDescriptorsKeepsStaleRelayHintWithinRetentionWindow(t *testing.T) {
34 - set, err := NewRelaySet(nil)
35 - if err != nil {
36 - t.Fatalf("NewRelaySet() error = %v", err)
37 - }
30 +func TestDescriptorsDropsExpiredSignedRelayDescriptor(t *testing.T) {
31 + set := NewRelaySet(nil)
32
33 now := time.Now().UTC()
34 relayURL := "https://relay-stale.example"
@@ -56,55 +50,14 @@ func TestDescriptorsKeepsStaleRelayHintWithinRetentionWindow(t *testing.T) {
50 set.relays[relayURL] = state
51 set.mu.Unlock()
52
59 - descriptors := set.Descriptors()
60 - if len(descriptors) != 1 {
61 - t.Fatalf("len(Descriptors()) = %d, want 1", len(descriptors))
62 - }
63 - got := descriptors[0]
64 - if got.APIHTTPSAddr != relayURL {
65 - t.Fatalf("descriptor api_https_addr = %q, want %q", got.APIHTTPSAddr, relayURL)
66 - }
67 - if !got.ExpiresAt.After(now) {
68 - t.Fatalf("descriptor expires_at = %v, want future expiry", got.ExpiresAt)
69 - }
70 - if !got.SupportsUDP || !got.SupportsTCP || !got.SupportsOverlayPeer {
71 - t.Fatal("stale advertised descriptor should preserve last known capability claims")
72 - }
73 - if got.IngressTLSAddr == "" || got.WireGuardPublicKey == "" || got.WireGuardEndpoint == "" || got.OverlayIPv4 == "" || len(got.OverlayCIDRs) == 0 {
74 - t.Fatal("stale advertised descriptor should preserve last known routing fields")
75 - }
76 - if got.Load != 1 || got.LoadScore != 2 {
77 - t.Fatal("stale advertised descriptor should preserve last known load signals")
78 - }
79 -}
80 -
81 -func TestDescriptorsDropsRelayAfterHintRetentionWindow(t *testing.T) {
82 - set, err := NewRelaySet(nil)
83 - if err != nil {
84 - t.Fatalf("NewRelaySet() error = %v", err)
85 - }
86 -
87 - now := time.Now().UTC()
88 - relayURL := "https://relay-old.example"
89 - state := confirmedPolicyRelayState(t, "relay-old", relayURL)
90 - state.Descriptor.ExpiresAt = now.Add(-time.Minute)
91 - state.LastSeenAt = now.Add(-DiscoveryHintRetentionTTL).Add(-time.Minute)
92 -
93 - set.mu.Lock()
94 - set.relays[relayURL] = state
95 - set.mu.Unlock()
96 -
97 - descriptors := set.Descriptors()
53 + descriptors := set.Descriptors(types.RelayDescriptor{})
54 if len(descriptors) != 0 {
99 - t.Fatalf("len(Descriptors()) = %d, want 0", len(descriptors))
55 + t.Fatalf("len(Descriptors(empty)) = %d, want 0", len(descriptors))
56 }
57 }
58
59 func TestApplyRelayDiscoveryResponseCollectsRelaysDespiteProtocolMismatch(t *testing.T) {
104 - set, err := NewRelaySet(nil)
105 - if err != nil {
106 - t.Fatalf("NewRelaySet() error = %v", err)
107 - }
60 + set := NewRelaySet(nil)
61
62 desc := mustPolicyRelayDescriptor(t, "relay-mismatch", "https://relay-mismatch.example")
63 changed, err := set.ApplyRelayDiscoveryResponse("", types.DiscoveryResponse{
@@ -131,10 +84,7 @@ func TestApplyRelayDiscoveryResponseCollectsRelaysDespiteProtocolMismatch(t *tes
84 }
85
86 func TestApplyRelayDiscoveryResponseCollectsHintsWhenTargetDescriptorIsMissing(t *testing.T) {
134 - set, err := NewRelaySet(nil)
135 - if err != nil {
136 - t.Fatalf("NewRelaySet() error = %v", err)
137 - }
87 + set := NewRelaySet(nil)
88
89 hinted := mustPolicyRelayDescriptor(t, "relay-hinted", "https://relay-hinted.example")
90 changed, err := set.ApplyRelayDiscoveryResponse("https://relay-source.example", types.DiscoveryResponse{
@@ -161,10 +111,7 @@ func TestApplyRelayDiscoveryResponseCollectsHintsWhenTargetDescriptorIsMissing(t
111 }
112
113 func TestApplyRelayDiscoveryResponseClearsDirectRetryOnAuthoritativeSuccess(t *testing.T) {
164 - set, err := NewRelaySet(nil)
165 - if err != nil {
166 - t.Fatalf("NewRelaySet() error = %v", err)
167 - }
114 + set := NewRelaySet(nil)
115
116 relayURL := "https://relay-source.example"
117 desc := mustPolicyRelayDescriptor(t, "relay-source", relayURL)
@@ -197,10 +144,7 @@ func TestApplyRelayDiscoveryResponseClearsDirectRetryOnAuthoritativeSuccess(t *t
144 }
145
146 func TestApplyRelayDiscoveryResponsePreservesDirectRetryOnHint(t *testing.T) {
200 - set, err := NewRelaySet(nil)
201 - if err != nil {
202 - t.Fatalf("NewRelaySet() error = %v", err)
203 - }
147 + set := NewRelaySet(nil)
148
149 relayURL := "https://relay-hinted.example"
150 desc := mustPolicyRelayDescriptor(t, "relay-hinted", relayURL)
@@ -230,10 +174,7 @@ func TestApplyRelayDiscoveryResponsePreservesDirectRetryOnHint(t *testing.T) {
174 }
175
176 func TestConfirmRelayURLMarksRelayConfirmedWithoutChangingAggregateDescriptor(t *testing.T) {
233 - set, err := NewRelaySet(nil)
234 - if err != nil {
235 - t.Fatalf("NewRelaySet() error = %v", err)
236 - }
177 + set := NewRelaySet(nil)
178
179 relayURL := "https://relay-confirmed.example"
180 state := RelayState{
@@ -259,10 +200,7 @@ func TestConfirmRelayURLMarksRelayConfirmedWithoutChangingAggregateDescriptor(t
200 }
201
202 func TestUnconfirmRelayURLClearsLocalConfirmationOnly(t *testing.T) {
262 - set, err := NewRelaySet(nil)
263 - if err != nil {
264 - t.Fatalf("NewRelaySet() error = %v", err)
265 - }
203 + set := NewRelaySet(nil)
204
205 relayURL := "https://relay-confirmed.example"
206 state := confirmedPolicyRelayState(t, "relay-confirmed", relayURL)
portal/discovery/relaystate.go
+10 -34
@@ -1,7 +1,6 @@
1 package discovery
2
3 import (
4 - "errors"
4 "time"
5
6 "github.com/gosuda/portal-tunnel/v2/types"
@@ -10,7 +9,6 @@ import (
9
10 const (
11 DiscoveryDescriptorTTL = 5 * time.Minute
13 - DiscoveryHintRetentionTTL = 30 * 24 * time.Hour
12 defaultDirectRecoveryBackoff = 1 * time.Minute
13 maxDirectRecoveryBackoff = 5 * time.Minute
14
@@ -46,37 +44,7 @@ type RelayState struct {
44 nextDirectRefreshAt time.Time
45 }
46
49 -type ClientState struct {
50 - ActiveRelayURLs []string
51 - ExplicitRelayURLs []string
52 - MaxActiveRelays int
53 - RequireUDP bool
54 - RequireTCP bool
55 -}
56 -
57 -func newRelayState(desc types.RelayDescriptor, seenAt time.Time) (RelayState, error) {
58 - state := RelayState{
59 - Descriptor: desc,
60 - }
61 - if seenAt.IsZero() {
62 - return state, nil
63 - }
64 -
65 - seenAt = seenAt.UTC()
66 - normalized, err := utils.NormalizeDescriptor(desc)
67 - if err != nil {
68 - return RelayState{}, err
69 - }
70 - if normalized.ExpiresAt.Before(seenAt) {
71 - return RelayState{}, errors.New("descriptor expired")
72 - }
73 -
74 - state.Descriptor = normalized
75 - state.LastSeenAt = seenAt
76 - return state, nil
77 -}
78 -
79 -func newRelayStateFromURL(relayURL string) RelayState {
47 +func newRelayState(relayURL string) RelayState {
48 return RelayState{
49 Descriptor: types.RelayDescriptor{
50 Identity: types.Identity{
@@ -88,6 +56,14 @@ func newRelayStateFromURL(relayURL string) RelayState {
56 }
57 }
58
91 -func (state RelayState) hasDescriptor() bool {
59 +func (state RelayState) hasObservedDescriptor() bool {
60 return !state.LastSeenAt.IsZero()
61 }
62 +
63 +type ClientState struct {
64 + ActiveRelayURLs []string
65 + ExplicitRelayURLs []string
66 + MaxActiveRelays int
67 + RequireUDP bool
68 + RequireTCP bool
69 +}
portal/discovery/sign.go deleted
-135
@@ -1,135 +0,0 @@
1 -package discovery
2 -
3 -import (
4 - "crypto/sha256"
5 - "encoding/base64"
6 - "encoding/hex"
7 - "errors"
8 - "fmt"
9 - "strings"
10 -
11 - "github.com/decred/dcrd/dcrec/secp256k1/v4/ecdsa"
12 -
13 - "github.com/gosuda/portal-tunnel/v2/types"
14 - "github.com/gosuda/portal-tunnel/v2/utils"
15 -)
16 -
17 -// descriptorSignatureSize is the byte length of a recoverable secp256k1
18 -// compact ECDSA signature: 1 byte recovery code + 32 byte R + 32 byte S.
19 -const descriptorSignatureSize = 65
20 -
21 -// ErrDescriptorUnsigned is returned when a descriptor that should carry a
22 -// signature is missing one. Distinct from ErrDescriptorInvalidSignature so
23 -// callers can differentiate "this peer never signed" from "this peer signed
24 -// incorrectly or the payload was tampered with".
25 -var (
26 - ErrDescriptorUnsigned = errors.New("relay descriptor is not signed")
27 - ErrDescriptorInvalidSignature = errors.New("relay descriptor signature is invalid")
28 - ErrDescriptorAddressMismatch = errors.New("relay descriptor address does not match recovered signing key")
29 - ErrDescriptorMissingAddress = errors.New("relay descriptor address is required for signature verification")
30 -)
31 -
32 -// SignDescriptor returns a copy of desc with its Signature field populated by
33 -// signing the canonical bytes with the supplied secp256k1 private key (hex
34 -// encoded). The signature is recoverable, so verifiers do not need to know
35 -// the public key out of band — they recover it from the signature and check
36 -// it derives the descriptor's Address field.
37 -//
38 -// Mutable telemetry fields (Load, LoadScore, LastUpdated) are NOT covered by
39 -// the signature, so callers may freely update them after signing without
40 -// invalidating the signature.
41 -func SignDescriptor(desc types.RelayDescriptor, privateKeyHex string) (types.RelayDescriptor, error) {
42 - privateKey, _, err := utils.ParseSecp256k1PrivateKeyHex(privateKeyHex, true)
43 - if err != nil {
44 - return types.RelayDescriptor{}, fmt.Errorf("relay descriptor signing key: %w", err)
45 - }
46 -
47 - // Strip any pre-existing signature so re-signing is idempotent and the
48 - // canonical bytes never depend on what the signature happens to be.
49 - desc.Signature = ""
50 -
51 - // Both signer and verifier hash the post-normalization form so that
52 - // JSON round-trips and minor input variation (case, whitespace, etc.)
53 - // do not break signature verification.
54 - normalized, err := utils.NormalizeDescriptor(desc)
55 - if err != nil {
56 - return types.RelayDescriptor{}, fmt.Errorf("normalize relay descriptor for signing: %w", err)
57 - }
58 - if strings.TrimSpace(normalized.Address) == "" {
59 - return types.RelayDescriptor{}, ErrDescriptorMissingAddress
60 - }
61 - desc = normalized
62 -
63 - canonical, err := types.CanonicalBytes(desc)
64 - if err != nil {
65 - return types.RelayDescriptor{}, fmt.Errorf("canonicalize relay descriptor: %w", err)
66 - }
67 - digest := sha256.Sum256(canonical)
68 - // isCompressedKey=true matches the descriptor's Identity.PublicKey, which
69 - // is the compressed form of the public key throughout the codebase.
70 - signature := ecdsa.SignCompact(privateKey, digest[:], true)
71 - if len(signature) != descriptorSignatureSize {
72 - return types.RelayDescriptor{}, fmt.Errorf("unexpected compact signature length %d", len(signature))
73 - }
74 -
75 - desc.Signature = base64.StdEncoding.EncodeToString(signature)
76 - return desc, nil
77 -}
78 -
79 -// VerifyDescriptor checks the descriptor's signature against its canonical
80 -// bytes and confirms that the recovered signing key corresponds to the
81 -// descriptor's Address field. Returns the verified compressed public key hex
82 -// on success so callers can use it as a stable identity for indexing.
83 -//
84 -// The descriptor argument is treated as read-only.
85 -func VerifyDescriptor(desc types.RelayDescriptor) (publicKeyHex string, err error) {
86 - rawSignature := strings.TrimSpace(desc.Signature)
87 - if rawSignature == "" {
88 - return "", ErrDescriptorUnsigned
89 - }
90 -
91 - signature, err := base64.StdEncoding.DecodeString(rawSignature)
92 - if err != nil {
93 - return "", fmt.Errorf("%w: base64 decode: %w", ErrDescriptorInvalidSignature, err)
94 - }
95 - if len(signature) != descriptorSignatureSize {
96 - return "", fmt.Errorf("%w: unexpected signature length %d", ErrDescriptorInvalidSignature, len(signature))
97 - }
98 -
99 - // Strip the signature and normalize before recomputing canonical bytes
100 - // so signer and verifier agree on the exact byte sequence that was
101 - // hashed regardless of incidental input variation (case, whitespace,
102 - // JSON ordering of slice elements normalized in NormalizeDescriptor).
103 - unsignedCopy := desc
104 - unsignedCopy.Signature = ""
105 - normalized, err := utils.NormalizeDescriptor(unsignedCopy)
106 - if err != nil {
107 - return "", fmt.Errorf("%w: normalize: %w", ErrDescriptorInvalidSignature, err)
108 - }
109 - if strings.TrimSpace(normalized.Address) == "" {
110 - return "", ErrDescriptorMissingAddress
111 - }
112 - canonical, err := types.CanonicalBytes(normalized)
113 - if err != nil {
114 - return "", fmt.Errorf("canonicalize relay descriptor: %w", err)
115 - }
116 - digest := sha256.Sum256(canonical)
117 -
118 - publicKey, _, err := ecdsa.RecoverCompact(signature, digest[:])
119 - if err != nil {
120 - return "", fmt.Errorf("%w: %w", ErrDescriptorInvalidSignature, err)
121 - }
122 - if publicKey == nil {
123 - return "", ErrDescriptorInvalidSignature
124 - }
125 -
126 - publicKeyHex = hex.EncodeToString(publicKey.SerializeCompressed())
127 - derivedAddress, err := utils.AddressFromCompressedPublicKeyHex(publicKeyHex)
128 - if err != nil {
129 - return "", fmt.Errorf("derive address from recovered key: %w", err)
130 - }
131 - if !strings.EqualFold(strings.TrimSpace(derivedAddress), strings.TrimSpace(normalized.Address)) {
132 - return "", ErrDescriptorAddressMismatch
133 - }
134 - return publicKeyHex, nil
135 -}
portal/discovery/sign_test.go deleted
-210
@@ -1,210 +0,0 @@
1 -package discovery
2 -
3 -import (
4 - "errors"
5 - "testing"
6 - "time"
7 -
8 - "github.com/gosuda/portal-tunnel/v2/types"
9 - "github.com/gosuda/portal-tunnel/v2/utils"
10 -)
11 -
12 -func mustSigningIdentity(t *testing.T) types.Identity {
13 - t.Helper()
14 - identity, err := utils.ResolveSecp256k1Identity("")
15 - if err != nil {
16 - t.Fatalf("ResolveSecp256k1Identity() error = %v", err)
17 - }
18 - return identity
19 -}
20 -
21 -func mustNormalizedDescriptor(t *testing.T, signing types.Identity, relayName, relayURL string) types.RelayDescriptor {
22 - t.Helper()
23 - now := time.Now().UTC().Truncate(time.Microsecond)
24 - desc, err := utils.NormalizeDescriptor(types.RelayDescriptor{
25 - Identity: types.Identity{
26 - Name: relayName,
27 - Address: signing.Address,
28 - },
29 - RelayID: relayURL,
30 - Version: 1,
31 - IssuedAt: now,
32 - ExpiresAt: now.Add(time.Hour),
33 - APIHTTPSAddr: relayURL,
34 - Discovery: true,
35 - })
36 - if err != nil {
37 - t.Fatalf("NormalizeDescriptor() error = %v", err)
38 - }
39 - return desc
40 -}
41 -
42 -func TestSignDescriptorRoundtrip(t *testing.T) {
43 - signing := mustSigningIdentity(t)
44 - desc := mustNormalizedDescriptor(t, signing, "relay-rt", "https://relay-rt.example")
45 - signed, err := SignDescriptor(desc, signing.PrivateKey)
46 - if err != nil {
47 - t.Fatalf("SignDescriptor() error = %v", err)
48 - }
49 - if signed.Signature == "" {
50 - t.Fatal("signed descriptor must have non-empty Signature")
51 - }
52 - pubKey, err := VerifyDescriptor(signed)
53 - if err != nil {
54 - t.Fatalf("VerifyDescriptor() error = %v", err)
55 - }
56 - if pubKey == "" {
57 - t.Fatal("VerifyDescriptor must return recovered public key")
58 - }
59 -}
60 -
61 -func TestVerifyDescriptorRejectsUnsigned(t *testing.T) {
62 - signing := mustSigningIdentity(t)
63 - desc := mustNormalizedDescriptor(t, signing, "relay-unsigned", "https://relay-unsigned.example")
64 - if _, err := VerifyDescriptor(desc); !errors.Is(err, ErrDescriptorUnsigned) {
65 - t.Fatalf("VerifyDescriptor() err = %v, want ErrDescriptorUnsigned", err)
66 - }
67 -}
68 -
69 -func TestVerifyDescriptorRejectsTamperedSignedField(t *testing.T) {
70 - signing := mustSigningIdentity(t)
71 - desc := mustNormalizedDescriptor(t, signing, "relay-tamper", "https://relay-tamper.example")
72 - signed, err := SignDescriptor(desc, signing.PrivateKey)
73 - if err != nil {
74 - t.Fatalf("SignDescriptor() error = %v", err)
75 - }
76 - tampered := signed
77 - tampered.WireGuardEndpoint = "evil.example:51820"
78 - if _, err := VerifyDescriptor(tampered); err == nil {
79 - t.Fatal("VerifyDescriptor must reject tampered signed field")
80 - }
81 -}
82 -
83 -func TestVerifyDescriptorAcceptsTelemetryUpdate(t *testing.T) {
84 - signing := mustSigningIdentity(t)
85 - desc := mustNormalizedDescriptor(t, signing, "relay-telemetry", "https://relay-telemetry.example")
86 - signed, err := SignDescriptor(desc, signing.PrivateKey)
87 - if err != nil {
88 - t.Fatalf("SignDescriptor() error = %v", err)
89 - }
90 - updated := signed
91 - updated.Load = 42
92 - updated.LoadScore = 99
93 - updated.LastUpdated = time.Now().UnixMilli()
94 - if _, err := VerifyDescriptor(updated); err != nil {
95 - t.Fatalf("VerifyDescriptor() must ignore telemetry, got err = %v", err)
96 - }
97 -}
98 -
99 -func TestVerifyDescriptorRejectsAddressMismatch(t *testing.T) {
100 - signing := mustSigningIdentity(t)
101 - other := mustSigningIdentity(t)
102 - desc := mustNormalizedDescriptor(t, signing, "relay-mismatch", "https://relay-mismatch.example")
103 - // Sign with `signing` but rewrite Address to a different identity.
104 - signed, err := SignDescriptor(desc, signing.PrivateKey)
105 - if err != nil {
106 - t.Fatalf("SignDescriptor() error = %v", err)
107 - }
108 - signed.Address = other.Address
109 - if _, err := VerifyDescriptor(signed); err == nil {
110 - t.Fatal("VerifyDescriptor must reject signing-key/address mismatch")
111 - }
112 -}
113 -
114 -func TestCanonicalBytesDeterministic(t *testing.T) {
115 - signing := mustSigningIdentity(t)
116 - desc := mustNormalizedDescriptor(t, signing, "relay-det", "https://relay-det.example")
117 - first, err := types.CanonicalBytes(desc)
118 - if err != nil {
119 - t.Fatalf("CanonicalBytes() error = %v", err)
120 - }
121 - for i := range 16 {
122 - out, err := types.CanonicalBytes(desc)
123 - if err != nil {
124 - t.Fatalf("CanonicalBytes() error = %v", err)
125 - }
126 - if string(out) != string(first) {
127 - t.Fatalf("CanonicalBytes is not deterministic: iteration %d differs", i)
128 - }
129 - }
130 -}
131 -
132 -func TestApplyRelayDiscoveryResponseRejectsUnsignedDescriptor(t *testing.T) {
133 - set, err := NewRelaySet(nil)
134 - if err != nil {
135 - t.Fatalf("NewRelaySet() error = %v", err)
136 - }
137 - signing := mustSigningIdentity(t)
138 - desc := mustNormalizedDescriptor(t, signing, "relay-strict", "https://relay-strict.example")
139 - // Note: NOT signed.
140 - _, err = set.ApplyRelayDiscoveryResponse("", types.DiscoveryResponse{
141 - ProtocolVersion: types.DiscoveryVersion,
142 - Relays: []types.RelayDescriptor{desc},
143 - }, time.Now().UTC())
144 - if err != nil {
145 - t.Fatalf("ApplyRelayDiscoveryResponse() error = %v", err)
146 - }
147 - if got := set.AggregateRelays(); len(got) != 0 {
148 - t.Fatalf("expected unsigned descriptor to be dropped, got %d relays", len(got))
149 - }
150 -}
151 -
152 -func TestApplyRelayDiscoveryResponseRejectsRollback(t *testing.T) {
153 - set, err := NewRelaySet(nil)
154 - if err != nil {
155 - t.Fatalf("NewRelaySet() error = %v", err)
156 - }
157 - signing := mustSigningIdentity(t)
158 - relayURL := "https://relay-rollback.example"
159 -
160 - now := time.Now().UTC().Truncate(time.Microsecond)
161 - build := func(issuedAt time.Time) types.RelayDescriptor {
162 - desc, err := utils.NormalizeDescriptor(types.RelayDescriptor{
163 - Identity: types.Identity{
164 - Name: "relay-rollback",
165 - Address: signing.Address,
166 - },
167 - RelayID: relayURL,
168 - Version: 1,
169 - IssuedAt: issuedAt,
170 - ExpiresAt: issuedAt.Add(time.Hour),
171 - APIHTTPSAddr: relayURL,
172 - Discovery: true,
173 - })
174 - if err != nil {
175 - t.Fatalf("NormalizeDescriptor() error = %v", err)
176 - }
177 - signed, err := SignDescriptor(desc, signing.PrivateKey)
178 - if err != nil {
179 - t.Fatalf("SignDescriptor() error = %v", err)
180 - }
181 - return signed
182 - }
183 -
184 - newer := build(now)
185 - if _, err := set.ApplyRelayDiscoveryResponse("", types.DiscoveryResponse{
186 - ProtocolVersion: types.DiscoveryVersion,
187 - Relays: []types.RelayDescriptor{newer},
188 - }, now); err != nil {
189 - t.Fatalf("apply newer error = %v", err)
190 - }
191 -
192 - older := build(now.Add(-time.Minute))
193 - changed, err := set.ApplyRelayDiscoveryResponse("", types.DiscoveryResponse{
194 - ProtocolVersion: types.DiscoveryVersion,
195 - Relays: []types.RelayDescriptor{older},
196 - }, now)
197 - if err != nil {
198 - t.Fatalf("apply older error = %v", err)
199 - }
200 - if changed {
201 - t.Fatal("expected rollback descriptor to leave relay set unchanged")
202 - }
203 - states := set.AggregateRelays()
204 - if len(states) != 1 {
205 - t.Fatalf("len(AggregateRelays()) = %d, want 1", len(states))
206 - }
207 - if !states[0].Descriptor.IssuedAt.Equal(newer.IssuedAt) {
208 - t.Fatalf("retained descriptor IssuedAt = %v, want %v", states[0].Descriptor.IssuedAt, newer.IssuedAt)
209 - }
210 -}
portal/lease.go
+19 -29
@@ -204,7 +204,7 @@ func (r *leaseRegistry) issueRegisterChallenge(req types.RegisterChallengeReques
204 func (r *leaseRegistry) consumeVerifiedRegisterChallenge(req types.RegisterRequest) (*auth.RegisterChallenge, error) {
205 challengeID := strings.TrimSpace(req.ChallengeID)
206 if challengeID == "" {
207 - return nil, auth.ErrChallengeNotFound
207 + return nil, auth.ErrRegisterChallengeNotFound
208 }
209
210 now := time.Now().UTC()
@@ -213,11 +213,11 @@ func (r *leaseRegistry) consumeVerifiedRegisterChallenge(req types.RegisterReque
213
214 challenge := r.registerChallenges[challengeID]
215 if challenge == nil {
216 - return nil, auth.ErrChallengeNotFound
216 + return nil, auth.ErrRegisterChallengeNotFound
217 }
218 if challenge.Expired(now) {
219 delete(r.registerChallenges, challengeID)
220 - return nil, auth.ErrChallengeExpired
220 + return nil, auth.ErrRegisterChallengeExpired
221 }
222 if err := challenge.Verify(req, now); err != nil {
223 return nil, err
@@ -227,20 +227,19 @@ func (r *leaseRegistry) consumeVerifiedRegisterChallenge(req types.RegisterReque
227 return challenge, nil
228 }
229
230 -func (r *leaseRegistry) Touch(identity types.Identity, clientIP string, now time.Time) *leaseRecord {
230 +func (r *leaseRegistry) Touch(identity types.Identity, clientIP string, now time.Time) {
231 r.mu.Lock()
232 defer r.mu.Unlock()
233
234 record, ok := r.leasesByKey[identity.Key()]
235 if !ok {
236 - return nil
236 + return
237 }
238 record.LastSeenAt = now
239 if strings.TrimSpace(clientIP) != "" {
240 record.ClientIP = clientIP
241 }
242 r.policy.IPFilter().RegisterIdentityIP(record.Key(), clientIP)
243 - return record
243 }
244
245 func (r *leaseRegistry) cleanupExpired(now time.Time) []*leaseRecord {
@@ -257,7 +256,7 @@ func (r *leaseRegistry) cleanupExpired(now time.Time) []*leaseRecord {
256 }
257 }
258 for challengeID, challenge := range r.registerChallenges {
260 - if challenge == nil || challenge.Expired(now) {
259 + if challenge.Expired(now) {
260 delete(r.registerChallenges, challengeID)
261 }
262 }
@@ -292,36 +291,27 @@ func (r *leaseRegistry) countTCPPortLeases() int {
291 return count
292 }
293
295 -func (r *leaseRegistry) snapshot(record *leaseRecord, now time.Time) (types.Lease, bool) {
296 - if record == nil || now.After(record.ExpiresAt) {
297 - return types.Lease{}, false
298 - }
299 -
300 - adminSnapshot := r.AdminSnapshot(record)
301 - since := time.Duration(0)
302 - if !adminSnapshot.LastSeenAt.IsZero() {
303 - since = max(now.Sub(adminSnapshot.LastSeenAt), 0)
304 - }
305 - if adminSnapshot.IsBanned || adminSnapshot.IsDenied || !adminSnapshot.IsApproved || adminSnapshot.Metadata.Hide {
306 - return types.Lease{}, false
307 - }
308 - if adminSnapshot.Ready == 0 && since >= 3*time.Minute {
309 - return types.Lease{}, false
310 - }
311 - return adminSnapshot.Lease, true
312 -}
313 -
294 func (r *leaseRegistry) LeaseSnapshots(now time.Time) []types.Lease {
295 r.mu.RLock()
296 defer r.mu.RUnlock()
297
298 snapshots := make([]types.Lease, 0, len(r.leasesByKey))
299 for _, record := range r.leasesByKey {
320 - snapshot, ok := r.snapshot(record, now)
321 - if !ok {
300 + if record == nil || now.After(record.ExpiresAt) {
301 + continue
302 + }
303 + adminSnapshot := r.AdminSnapshot(record)
304 + since := time.Duration(0)
305 + if !adminSnapshot.LastSeenAt.IsZero() {
306 + since = max(now.Sub(adminSnapshot.LastSeenAt), 0)
307 + }
308 + if adminSnapshot.IsBanned || adminSnapshot.IsDenied || !adminSnapshot.IsApproved || adminSnapshot.Metadata.Hide {
309 + continue
310 + }
311 + if adminSnapshot.Ready == 0 && since >= 3*time.Minute {
312 continue
313 }
324 - snapshots = append(snapshots, snapshot)
314 + snapshots = append(snapshots, adminSnapshot.Lease)
315 }
316 return snapshots
317 }
portal/server.go
+7 -13
@@ -145,10 +145,7 @@ func NewServer(cfg ServerConfig) (*Server, error) {
145 return nil, fmt.Errorf("resolve discovery bootstraps: %w", err)
146 }
147 cfg.Bootstraps = utils.RemoveRelayURL(cfg.Bootstraps, cfg.PortalURL)
148 - relaySet, err = discovery.NewRelaySet(cfg.Bootstraps)
149 - if err != nil {
150 - return nil, err
151 - }
148 + relaySet = discovery.NewRelaySet(cfg.Bootstraps)
149 }
150
151 return &Server{
@@ -598,20 +595,17 @@ func (s *Server) runRelayDiscoveryLoop(ctx context.Context) error {
595 <-ctx.Done()
596 return nil
597 }
601 - refresher, err := discovery.NewRefresher(s.relaySet, nil, s.overlay, s.PortalURL())
602 - if err != nil {
603 - return err
604 - }
598 + refresher := discovery.NewRefresher(s.relaySet, s.overlay)
599 ticker := time.NewTicker(discovery.DiscoveryPollInterval)
600 defer ticker.Stop()
601
602 for {
609 - snapshots := s.registry.LeaseSnapshots(time.Now())
610 - sourceHosts := make([]string, 0, len(snapshots))
611 - for _, snapshot := range snapshots {
612 - sourceHosts = append(sourceHosts, snapshot.Hostname)
603 + now := time.Now().UTC()
604 + self, err := s.signedRelayDescriptor(now)
605 + if err != nil {
606 + return fmt.Errorf("build relay discovery descriptor: %w", err)
607 }
614 - if err := refresher.Refresh(ctx, sourceHosts...); err != nil {
608 + if err := refresher.Refresh(ctx, self); err != nil {
609 if ctx.Err() != nil {
610 return nil
611 }
sdk/expose.go
+1 -6
@@ -102,11 +102,6 @@ func Expose(ctx context.Context, cfg ExposeConfig) (*Exposure, error) {
102 return nil, fmt.Errorf("invalid --udp-addr value %q: %w", cfg.UDPAddr, err)
103 }
104 }
105 - relaySet, err := discovery.NewRelaySet(relayURLs)
106 - if err != nil {
107 - return nil, err
108 - }
109 -
105 exposureCtx, cancel := context.WithCancel(ctx)
106 exposure := &Exposure{
107 cancel: cancel,
@@ -123,7 +118,7 @@ func Expose(ctx context.Context, cfg ExposeConfig) (*Exposure, error) {
118 rootCAPEM: append([]byte(nil), cfg.RootCAPEM...),
119 accepted: make(chan net.Conn, max(len(relayURLs)*defaultReadyTarget*2, 1)),
120 datagrams: make(chan types.DatagramFrame, max(len(relayURLs)*32, 1)),
126 - relaySet: relaySet,
121 + relaySet: discovery.NewRelaySet(relayURLs),
122 relayListeners: make(map[string]*Listener, len(relayURLs)),
123 }
124
sdk/expose_test.go
+2 -9
@@ -9,12 +9,7 @@ import (
9
10 func mustRelaySet(t *testing.T, relayURLs ...string) *discovery.RelaySet {
11 t.Helper()
12 -
13 - set, err := discovery.NewRelaySet(relayURLs)
14 - if err != nil {
15 - t.Fatalf("NewRelaySet() error = %v", err)
16 - }
17 - return set
12 + return discovery.NewRelaySet(relayURLs)
13 }
14
15 func TestExposureReconcileRemovesBannedRelayFromActiveSet(t *testing.T) {
@@ -102,9 +97,7 @@ func TestExposureReconcileRemovesStaleListener(t *testing.T) {
97 },
98 }
99
105 - if err := exposure.relaySet.SetBootstrapRelayURLs([]string{relayB}); err != nil {
106 - t.Fatalf("SetBootstrapRelayURLs() error = %v", err)
107 - }
100 + exposure.relaySet.SetBootstrapRelayURLs([]string{relayB})
101 if err := exposure.reconcileRelayListeners(false); err != nil {
102 t.Fatalf("reconcileRelayListeners() error = %v", err)
103 }
utils/crypto.go
+92 -2
@@ -21,6 +21,19 @@ import (
21 "github.com/gosuda/portal-tunnel/v2/types"
22 )
23
24 +const (
25 + // CompactSecp256k1SignatureSize is the byte length of a compact
26 + // recoverable secp256k1 ECDSA signature.
27 + CompactSecp256k1SignatureSize = 65
28 + // RawSecp256k1SignatureSize is the byte length of the JOSE ES256K
29 + // signature form, r || s with no recovery header.
30 + RawSecp256k1SignatureSize = 64
31 +)
32 +
33 +// ErrSecp256k1SignatureInvalid marks a well-formed signature that does not
34 +// verify for the payload and public key.
35 +var ErrSecp256k1SignatureInvalid = errors.New("signature is invalid")
36 +
37 func NormalizeEVMAddress(raw string) (string, error) {
38 trimmed := strings.TrimSpace(raw)
39 if trimmed == "" {
@@ -152,6 +165,73 @@ func SignSHA256Secp256k1DER(payload []byte, privateKeyHex string) (string, error
165 return hex.EncodeToString(signature.Serialize()), nil
166 }
167
168 +// SignSHA256Secp256k1Compact signs the SHA-256 digest of payload and returns
169 +// the compact recoverable secp256k1 ECDSA signature.
170 +func SignSHA256Secp256k1Compact(payload []byte, privateKey *secp256k1.PrivateKey, compressed bool) ([]byte, error) {
171 + if privateKey == nil {
172 + return nil, errors.New("signing key is required")
173 + }
174 + hash := sha256.Sum256(payload)
175 + signature := ecdsa.SignCompact(privateKey, hash[:], compressed)
176 + if len(signature) != CompactSecp256k1SignatureSize {
177 + return nil, errors.New("invalid compact signature length")
178 + }
179 + return signature, nil
180 +}
181 +
182 +// SignSHA256Secp256k1Raw64 signs the SHA-256 digest of payload and returns
183 +// the raw r || s signature form used by ES256K.
184 +func SignSHA256Secp256k1Raw64(payload []byte, privateKey *secp256k1.PrivateKey) ([]byte, error) {
185 + compact, err := SignSHA256Secp256k1Compact(payload, privateKey, false)
186 + if err != nil {
187 + return nil, err
188 + }
189 +
190 + signature := make([]byte, RawSecp256k1SignatureSize)
191 + copy(signature[:32], compact[1:33])
192 + copy(signature[32:], compact[33:65])
193 + return signature, nil
194 +}
195 +
196 +// RecoverSHA256Secp256k1Compact recovers the public key from a compact
197 +// recoverable signature over the SHA-256 digest of payload.
198 +func RecoverSHA256Secp256k1Compact(payload, signature []byte) (*secp256k1.PublicKey, error) {
199 + if len(signature) != CompactSecp256k1SignatureSize {
200 + return nil, errors.New("invalid compact signature length")
201 + }
202 +
203 + hash := sha256.Sum256(payload)
204 + publicKey, _, err := ecdsa.RecoverCompact(signature, hash[:])
205 + if err != nil {
206 + return nil, err
207 + }
208 + if publicKey == nil {
209 + return nil, ErrSecp256k1SignatureInvalid
210 + }
211 + return publicKey, nil
212 +}
213 +
214 +// VerifySHA256Secp256k1Raw64 verifies an ES256K raw r || s signature over the
215 +// SHA-256 digest of payload.
216 +func VerifySHA256Secp256k1Raw64(payload, signature []byte, publicKey *secp256k1.PublicKey) error {
217 + if publicKey == nil {
218 + return errors.New("verification key is required")
219 + }
220 + if len(signature) != RawSecp256k1SignatureSize {
221 + return errors.New("invalid es256k signature length")
222 + }
223 +
224 + var r, s secp256k1.ModNScalar
225 + if overflow := r.SetByteSlice(signature[:32]); overflow || r.IsZero() {
226 + return errors.New("invalid es256k signature r")
227 + }
228 + if overflow := s.SetByteSlice(signature[32:]); overflow || s.IsZero() {
229 + return errors.New("invalid es256k signature s")
230 + }
231 +
232 + return verifySHA256Secp256k1Signature(payload, ecdsa.NewSignature(&r, &s), publicKey)
233 +}
234 +
235 func VerifySHA256Secp256k1DER(payload []byte, publicKeyHex, signatureHex string) error {
236 pubKey, err := ParseSecp256k1PublicKeyHex(publicKeyHex)
237 if err != nil {
@@ -173,9 +253,19 @@ func VerifySHA256Secp256k1DER(payload []byte, publicKeyHex, signatureHex string)
253 return fmt.Errorf("parse signature: %w", err)
254 }
255
256 + return verifySHA256Secp256k1Signature(payload, signature, pubKey)
257 +}
258 +
259 +func verifySHA256Secp256k1Signature(payload []byte, signature *ecdsa.Signature, publicKey *secp256k1.PublicKey) error {
260 + if signature == nil {
261 + return errors.New("signature is required")
262 + }
263 + if publicKey == nil {
264 + return errors.New("verification key is required")
265 + }
266 hash := sha256.Sum256(payload)
177 - if !signature.Verify(hash[:], pubKey) {
178 - return errors.New("signature is invalid")
267 + if !signature.Verify(hash[:], publicKey) {
268 + return ErrSecp256k1SignatureInvalid
269 }
270 return nil
271 }