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