fix: update session token generation and enhance register challenge handling
Kim committed
Apr 30, 2026 at 14:48 UTC
6f6a2584e6fedc94897c11952a7f17fd636a9e06
6 files changed
+131
-115
cmd/relay-server/admin.go
+1
-4
@@ -49,10 +49,7 @@ func (a *adminAuth) ValidateKey(key string) bool {
49
}
50
51
func (a *adminAuth) CreateSession() (string, error) {
52
- token, err := utils.RandomHex(32)
53
- if err != nil {
54
- return "", err
55
- }
52
+ token := utils.RandomID("")
53
54
a.mu.Lock()
55
defer a.mu.Unlock()
portal/api_server.go
+17
-15
@@ -29,19 +29,20 @@ type apiError struct {
29
func (e *apiError) Error() string { return e.msg }
30
31
var (
32
- errFeatureUnavailable = &apiError{types.APIErrorCodeFeatureUnavailable, "feature unavailable", http.StatusServiceUnavailable}
33
- errHostnameConflict = &apiError{types.APIErrorCodeHostnameConflict, "hostname conflict", http.StatusConflict}
34
- errIPBanned = &apiError{types.APIErrorCodeIPBanned, "request denied because source IP is banned", http.StatusForbidden}
35
- errLeaseNotFound = &apiError{types.APIErrorCodeLeaseNotFound, "lease not found", http.StatusNotFound}
36
- errLeaseRejected = &apiError{types.APIErrorCodeLeaseRejected, "lease is not approved for routing", http.StatusForbidden}
37
- errTransportMismatch = &apiError{types.APIErrorCodeTransportMismatch, "transport mismatch", http.StatusConflict}
38
- errUnauthorized = &apiError{types.APIErrorCodeUnauthorized, "unauthorized", http.StatusForbidden}
39
- errUDPDisabled = &apiError{types.APIErrorCodeUDPDisabled, "udp disabled", http.StatusForbidden}
40
- errUDPCapacityExceeded = &apiError{types.APIErrorCodeUDPCapacityExceeded, "udp capacity exceeded", http.StatusServiceUnavailable}
41
- errUDPPortExhausted = &apiError{types.APIErrorCodeUDPPortExhausted, "no udp ports available", http.StatusServiceUnavailable}
42
- errTCPPortDisabled = &apiError{types.APIErrorCodeTCPPortDisabled, "tcp port disabled", http.StatusForbidden}
43
- errTCPPortCapacityExceeded = &apiError{types.APIErrorCodeTCPPortCapacityExceeded, "tcp port capacity exceeded", http.StatusServiceUnavailable}
44
- errTCPPortExhausted = &apiError{types.APIErrorCodeTCPPortExhausted, "no tcp ports available", http.StatusServiceUnavailable}
32
+ errFeatureUnavailable = &apiError{types.APIErrorCodeFeatureUnavailable, "feature unavailable", http.StatusServiceUnavailable}
33
+ errHostnameConflict = &apiError{types.APIErrorCodeHostnameConflict, "hostname conflict", http.StatusConflict}
34
+ errIPBanned = &apiError{types.APIErrorCodeIPBanned, "request denied because source IP is banned", http.StatusForbidden}
35
+ errLeaseNotFound = &apiError{types.APIErrorCodeLeaseNotFound, "lease not found", http.StatusNotFound}
36
+ errLeaseRejected = &apiError{types.APIErrorCodeLeaseRejected, "lease is not approved for routing", http.StatusForbidden}
37
+ errTransportMismatch = &apiError{types.APIErrorCodeTransportMismatch, "transport mismatch", http.StatusConflict}
38
+ errUnauthorized = &apiError{types.APIErrorCodeUnauthorized, "unauthorized", http.StatusForbidden}
39
+ errUDPDisabled = &apiError{types.APIErrorCodeUDPDisabled, "udp disabled", http.StatusForbidden}
40
+ errUDPCapacityExceeded = &apiError{types.APIErrorCodeUDPCapacityExceeded, "udp capacity exceeded", http.StatusServiceUnavailable}
41
+ errUDPPortExhausted = &apiError{types.APIErrorCodeUDPPortExhausted, "no udp ports available", http.StatusServiceUnavailable}
42
+ errTCPPortDisabled = &apiError{types.APIErrorCodeTCPPortDisabled, "tcp port disabled", http.StatusForbidden}
43
+ errTCPPortCapacityExceeded = &apiError{types.APIErrorCodeTCPPortCapacityExceeded, "tcp port capacity exceeded", http.StatusServiceUnavailable}
44
+ errTCPPortExhausted = &apiError{types.APIErrorCodeTCPPortExhausted, "no tcp ports available", http.StatusServiceUnavailable}
45
+ errRegisterChallengePending = &apiError{types.APIErrorCodeRateLimited, "too many pending register challenges", http.StatusTooManyRequests}
46
)
47
48
func writeAPIErrorResponse(w http.ResponseWriter, err error) {
@@ -296,7 +297,8 @@ func (s *Server) handleRegisterChallenge(w http.ResponseWriter, r *http.Request)
297
return
298
}
299
299
- if _, ok := s.extractAllowedClientIP(w, r); !ok {
300
+ clientIP, ok := s.extractAllowedClientIP(w, r)
301
+ if !ok {
302
return
303
}
304
@@ -332,7 +334,7 @@ func (s *Server) handleRegisterChallenge(w http.ResponseWriter, r *http.Request)
334
return
335
}
336
335
- resp, err := s.registry.issueRegisterChallenge(req, domain, registerURI)
337
+ resp, err := s.registry.issueRegisterChallenge(req, domain, registerURI, clientIP)
338
if err != nil {
339
writeAPIErrorResponse(w, err)
340
return
portal/lease.go
+68
-61
@@ -18,11 +18,12 @@ import (
18
)
19
20
const (
21
- defaultLeaseTTL = 30 * time.Second
22
- defaultRegisterChallengeTTL = 2 * time.Minute
23
- defaultPortReservationGrace = 5 * time.Minute
24
- defaultIdleKeepalive = 15 * time.Second
25
- defaultReadyQueueLimit = 8
21
+ defaultLeaseTTL = 30 * time.Second
22
+ defaultRegisterChallengeTTL = 2 * time.Minute
23
+ defaultRegisterChallengeOutstandingPerIP = 32
24
+ defaultPortReservationGrace = 5 * time.Minute
25
+ defaultIdleKeepalive = 15 * time.Second
26
+ defaultReadyQueueLimit = 8
27
)
28
29
type leaseRegistry struct {
@@ -161,17 +162,11 @@ func (r *leaseRegistry) Register(req types.RegisterChallengeRequest, clientIP, r
162
if !r.policy.IsUDPEnabled() {
163
return nil, types.RegisterResponse{}, errUDPDisabled
164
}
164
- if max := r.policy.UDPMaxLeases(); max > 0 && r.countDatagramLeases() >= max {
165
- return nil, types.RegisterResponse{}, errUDPCapacityExceeded
166
- }
165
}
166
if req.TCPEnabled {
167
if !r.policy.IsTCPPortEnabled() {
168
return nil, types.RegisterResponse{}, errTCPPortDisabled
169
}
172
- if max := r.policy.TCPPortMaxLeases(); max > 0 && r.countTCPPortLeases() >= max {
173
- return nil, types.RegisterResponse{}, errTCPPortCapacityExceeded
174
- }
170
if r.proxy == nil {
171
return nil, types.RegisterResponse{}, errors.New("tcp proxy is not available")
172
}
@@ -241,23 +236,53 @@ func (r *leaseRegistry) Register(req types.RegisterChallengeRequest, clientIP, r
236
replacedIndex := -1
237
r.mu.Lock()
238
now := time.Now()
244
- for _, existing := range r.records {
245
- if existing == nil || !existing.isPublicEntry() || existing.isExpired(now) {
239
+ udpLeases := 0
240
+ tcpLeases := 0
241
+ for i, existing := range r.records {
242
+ if existing == nil {
243
continue
244
}
248
- if existing.Hostname == hostname && existing.Key() != identityKey {
245
+ existingKey := existing.Key()
246
+ if replacedIndex < 0 && existing.stream != nil && existingKey == identityKey {
247
+ replaced = existing
248
+ replacedIndex = i
249
+ }
250
+ if existing.isExpired(now) {
251
+ continue
252
+ }
253
+ if existingKey != identityKey {
254
+ if existing.datagram != nil {
255
+ udpLeases++
256
+ }
257
+ if existing.tcpPort != nil {
258
+ tcpLeases++
259
+ }
260
+ }
261
+ if existing.isPublicEntry() && existing.Hostname == hostname && existingKey != identityKey {
262
r.mu.Unlock()
263
record.Close()
264
return nil, types.RegisterResponse{}, errHostnameConflict
265
}
253
- }
254
- if hopToken != "" {
255
- if existing := r.recordByHopToken(hopToken, now); existing != nil && existing.Key() != identityKey {
266
+ if hopToken != "" && (existing.isHopMiddle() || existing.isHopExit()) && existing.hopToken == hopToken && existingKey != identityKey {
267
r.mu.Unlock()
268
record.Close()
269
return nil, types.RegisterResponse{}, errors.New("hop token conflict")
270
}
271
}
272
+ if record.datagram != nil {
273
+ if max := r.policy.UDPMaxLeases(); max > 0 && udpLeases >= max {
274
+ r.mu.Unlock()
275
+ record.Close()
276
+ return nil, types.RegisterResponse{}, errUDPCapacityExceeded
277
+ }
278
+ }
279
+ if record.tcpPort != nil {
280
+ if max := r.policy.TCPPortMaxLeases(); max > 0 && tcpLeases >= max {
281
+ r.mu.Unlock()
282
+ record.Close()
283
+ return nil, types.RegisterResponse{}, errTCPPortCapacityExceeded
284
+ }
285
+ }
286
for i := 0; i < len(r.records); i++ {
287
existing := r.records[i]
288
if existing != nil && existing.stream == nil && existing.isPublicEntry() &&
@@ -266,15 +291,8 @@ func (r *leaseRegistry) Register(req types.RegisterChallengeRequest, clientIP, r
291
i--
292
}
293
}
269
- for i, existing := range r.records {
270
- if existing != nil && existing.stream != nil && existing.Key() == identityKey {
271
- replaced = existing
272
- replacedIndex = i
273
- r.policy.ForgetIdentity(identityKey)
274
- break
275
- }
276
- }
294
if replacedIndex >= 0 {
295
+ r.policy.ForgetIdentity(identityKey)
296
r.records[replacedIndex] = record
297
} else {
298
r.records = append(r.records, record)
@@ -533,7 +551,7 @@ func (r *leaseRegistry) DeleteHopRoute(route *types.HopRoute) *leaseRecord {
551
return deleted
552
}
553
536
-func (r *leaseRegistry) issueRegisterChallenge(req types.RegisterChallengeRequest, domain, uri string) (types.RegisterChallengeResponse, error) {
554
+func (r *leaseRegistry) issueRegisterChallenge(req types.RegisterChallengeRequest, domain, uri, clientIP string) (types.RegisterChallengeResponse, error) {
555
if strings.TrimSpace(req.HopToken) != "" && (req.UDPEnabled || req.TCPEnabled) {
556
return types.RegisterChallengeResponse{}, errTransportMismatch
557
}
@@ -541,17 +559,11 @@ func (r *leaseRegistry) issueRegisterChallenge(req types.RegisterChallengeReques
559
if !r.policy.IsUDPEnabled() {
560
return types.RegisterChallengeResponse{}, errUDPDisabled
561
}
544
- if max := r.policy.UDPMaxLeases(); max > 0 && r.countDatagramLeases() >= max {
545
- return types.RegisterChallengeResponse{}, errUDPCapacityExceeded
546
- }
562
}
563
if req.TCPEnabled {
564
if !r.policy.IsTCPPortEnabled() {
565
return types.RegisterChallengeResponse{}, errTCPPortDisabled
566
}
552
- if max := r.policy.TCPPortMaxLeases(); max > 0 && r.countTCPPortLeases() >= max {
553
- return types.RegisterChallengeResponse{}, errTCPPortCapacityExceeded
554
- }
567
}
568
569
now := time.Now().UTC()
@@ -559,13 +571,36 @@ func (r *leaseRegistry) issueRegisterChallenge(req types.RegisterChallengeReques
571
if err != nil {
572
return types.RegisterChallengeResponse{}, err
573
}
574
+ clientIP = strings.ToLower(strings.TrimSpace(clientIP))
575
+ if clientIP == "" {
576
+ clientIP = "<unknown>"
577
+ }
578
579
r.mu.Lock()
580
+ defer r.mu.Unlock()
581
+
582
+ pending := 0
583
+ for i := 0; i < len(r.records); {
584
+ record := r.records[i]
585
+ if record != nil && record.registerChallenge != nil {
586
+ if record.isExpired(now) {
587
+ r.deleteRecord(i)
588
+ continue
589
+ }
590
+ if record.ClientIP == clientIP {
591
+ pending++
592
+ }
593
+ }
594
+ i++
595
+ }
596
+ if pending >= defaultRegisterChallengeOutstandingPerIP {
597
+ return types.RegisterChallengeResponse{}, errRegisterChallengePending
598
+ }
599
r.records = append(r.records, &leaseRecord{
600
ExpiresAt: challenge.ExpiresAt,
601
+ ClientIP: clientIP,
602
registerChallenge: challenge,
603
})
568
- r.mu.Unlock()
604
605
return types.RegisterChallengeResponse{
606
ChallengeID: challenge.ChallengeID,
@@ -642,34 +677,6 @@ func (r *leaseRegistry) cleanupExpired(now time.Time) []*leaseRecord {
677
return expired
678
}
679
645
-func (r *leaseRegistry) countDatagramLeases() int {
646
- r.mu.RLock()
647
- defer r.mu.RUnlock()
648
-
649
- now := time.Now()
650
- count := 0
651
- for _, record := range r.records {
652
- if record != nil && !record.isExpired(now) && record.datagram != nil {
653
- count++
654
- }
655
- }
656
- return count
657
-}
658
-
659
-func (r *leaseRegistry) countTCPPortLeases() int {
660
- r.mu.RLock()
661
- defer r.mu.RUnlock()
662
-
663
- now := time.Now()
664
- count := 0
665
- for _, record := range r.records {
666
- if record != nil && !record.isExpired(now) && record.tcpPort != nil {
667
- count++
668
- }
669
- }
670
- return count
671
-}
672
-
680
func (r *leaseRegistry) PublicLeases(now time.Time) []types.Lease {
681
r.mu.RLock()
682
defer r.mu.RUnlock()
portal/lease_test.go
+41
@@ -3,6 +3,7 @@ package portal
3
import (
4
"context"
5
"errors"
6
+ "fmt"
7
"net"
8
"testing"
9
"time"
@@ -209,6 +210,46 @@ func TestLeaseRegistryCleanupExpiredClosesBroker(t *testing.T) {
210
}
211
}
212
213
+func TestIssueRegisterChallengeBoundsPendingPerIP(t *testing.T) {
214
+ t.Parallel()
215
+
216
+ registry := newTestRegistry(t)
217
+ clientIP := "203.0.113.50"
218
+ for i := 0; i < defaultRegisterChallengeOutstandingPerIP; i++ {
219
+ _, err := registry.issueRegisterChallenge(types.RegisterChallengeRequest{
220
+ Identity: newTestLeaseIdentity(t, fmt.Sprintf("demo-%d", i)),
221
+ }, "example.com", "https://example.com/sdk/register", clientIP)
222
+ if err != nil {
223
+ t.Fatalf("issueRegisterChallenge(%d) error = %v", i, err)
224
+ }
225
+ }
226
+
227
+ _, err := registry.issueRegisterChallenge(types.RegisterChallengeRequest{
228
+ Identity: newTestLeaseIdentity(t, "overflow"),
229
+ }, "example.com", "https://example.com/sdk/register", clientIP)
230
+ if !errors.Is(err, errRegisterChallengePending) {
231
+ t.Fatalf("issueRegisterChallenge() error = %v, want pending limit", err)
232
+ }
233
+
234
+ expiredAt := time.Now().Add(-time.Second)
235
+ registry.mu.Lock()
236
+ for _, record := range registry.records {
237
+ if record == nil || record.registerChallenge == nil {
238
+ continue
239
+ }
240
+ record.ExpiresAt = expiredAt
241
+ record.registerChallenge.ExpiresAt = expiredAt
242
+ }
243
+ registry.mu.Unlock()
244
+
245
+ _, err = registry.issueRegisterChallenge(types.RegisterChallengeRequest{
246
+ Identity: newTestLeaseIdentity(t, "after-cleanup"),
247
+ }, "example.com", "https://example.com/sdk/register", clientIP)
248
+ if err != nil {
249
+ t.Fatalf("issueRegisterChallenge() after expired cleanup error = %v", err)
250
+ }
251
+}
252
+
253
func TestServerRunRegistryJanitorRejectsNonPositiveInterval(t *testing.T) {
254
t.Parallel()
255
sdk/expose.go
+4
-4
@@ -369,15 +369,15 @@ func (e *Exposure) Accept() (net.Conn, error) {
369
connID := e.connSeq.Add(1)
370
log.Info().
371
Uint64("conn_id", connID).
372
- Str("local_addr", utils.AddrString(conn.LocalAddr())).
373
- Str("remote_addr", utils.AddrString(conn.RemoteAddr())).
372
+ Str("local_addr", conn.LocalAddr().String()).
373
+ Str("remote_addr", conn.RemoteAddr().String()).
374
Msg("exposure connection accepted")
375
376
return &exposureConn{
377
Conn: conn,
378
id: connID,
379
- localAddr: utils.AddrString(conn.LocalAddr()),
380
- remoteAddr: utils.AddrString(conn.RemoteAddr()),
379
+ localAddr: conn.LocalAddr().String(),
380
+ remoteAddr: conn.RemoteAddr().String(),
381
}, nil
382
}
383
}
utils/utils.go
-31
@@ -7,9 +7,7 @@ import (
7
"encoding/hex"
8
"errors"
9
"fmt"
10
- "io"
10
"net"
12
- "net/netip"
11
"net/url"
12
"path"
13
"strings"
@@ -474,13 +472,6 @@ func IsLocalRelayHost(host string) bool {
472
return strings.HasSuffix(host, ".localhost")
473
}
474
477
-func AddrString(addr net.Addr) string {
478
- if addr == nil {
479
- return ""
480
- }
481
- return addr.String()
482
-}
483
-
475
func ValidateIPv4(raw string) error {
476
ip := net.ParseIP(strings.TrimSpace(raw))
477
if ip == nil || ip.To4() == nil {
@@ -489,14 +480,6 @@ func ValidateIPv4(raw string) error {
480
return nil
481
}
482
492
-func RandomHex(size int) (string, error) {
493
- buf := make([]byte, size)
494
- if _, err := io.ReadFull(rand.Reader, buf); err != nil {
495
- return "", fmt.Errorf("read random bytes: %w", err)
496
- }
497
- return hex.EncodeToString(buf), nil
498
-}
499
-
483
func SleepOrDone(ctx context.Context, d time.Duration) bool {
484
timer := time.NewTimer(d)
485
defer timer.Stop()
@@ -516,20 +499,6 @@ func RandomID(prefix string) string {
499
return prefix + hex.EncodeToString(buf)
500
}
501
519
-func NormalizeIPPrefixes(inputs []string) []string {
520
- return normalizeUniqueStrings(inputs, func(input string) string {
521
- input = strings.TrimSpace(input)
522
- if input == "" {
523
- return ""
524
- }
525
- prefix, err := netip.ParsePrefix(input)
526
- if err != nil {
527
- return ""
528
- }
529
- return prefix.String()
530
- })
531
-}
532
-
502
func normalizeUniqueStrings(inputs []string, normalize func(string) string) []string {
503
if len(inputs) == 0 {
504
return nil