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
}