feat: enhance SNI handling, improve lease registration logic, and refactor related components
gosunuts committed
Feb 24, 2026 at 22:19 UTC
a49f58298c31e16f2092764b8ab7e9292fe3bb4e
11 files changed
+356
-228
cmd/relay-server/registry.go
+7
-57
@@ -4,13 +4,11 @@ import (
4
"crypto/subtle"
5
"encoding/json"
6
"fmt"
7
- "net"
7
"net/http"
8
"strings"
9
"time"
10
11
"github.com/rs/zerolog/log"
13
- "gosuda.org/portal/cmd/relay-server/manager"
12
"gosuda.org/portal/portal"
13
"gosuda.org/portal/portal/utils/sni"
14
"gosuda.org/portal/sdk"
@@ -69,13 +67,6 @@ func (r *SDKRegistry) HandleRegister(w http.ResponseWriter, req *http.Request) {
67
return
68
}
69
72
- if registerReq.Address == "" {
73
- writeJSON(w, sdk.RegisterResponse{
74
- Success: false,
75
- Message: "address is required",
76
- })
77
- return
78
- }
70
if strings.TrimSpace(registerReq.ReverseToken) == "" {
71
writeJSON(w, sdk.RegisterResponse{
72
Success: false,
@@ -84,20 +75,10 @@ func (r *SDKRegistry) HandleRegister(w http.ResponseWriter, req *http.Request) {
75
return
76
}
77
87
- resolvedAddr, err := resolveLeaseAddress(req, registerReq.Address)
88
- if err != nil {
89
- writeJSON(w, sdk.RegisterResponse{
90
- Success: false,
91
- Message: err.Error(),
92
- })
93
- return
94
- }
95
-
78
// Create lease
79
lease := &portal.Lease{
80
ID: registerReq.LeaseID,
81
Name: registerReq.Name,
100
- Address: resolvedAddr,
82
Metadata: registerReq.Metadata,
83
Expires: time.Now().Add(30 * time.Second),
84
TLSEnabled: registerReq.TLSEnabled,
@@ -113,7 +94,10 @@ func (r *SDKRegistry) HandleRegister(w http.ResponseWriter, req *http.Request) {
94
return
95
}
96
116
- if err := r.registerSNIRoute(registerReq.LeaseID, registerReq.Name, resolvedAddr); err != nil {
97
+ // Clear dropped state in case this is a re-registration after disconnect
98
+ r.server.GetReverseHub().ClearDropped(registerReq.LeaseID)
99
+
100
+ if err := r.registerSNIRoute(registerReq.LeaseID, registerReq.Name); err != nil {
101
// Keep lease and route state consistent on partial failure.
102
r.server.GetLeaseManager().DeleteLease(registerReq.LeaseID)
103
writeJSON(w, sdk.RegisterResponse{
@@ -126,8 +110,6 @@ func (r *SDKRegistry) HandleRegister(w http.ResponseWriter, req *http.Request) {
110
log.Info().
111
Str("lease_id", registerReq.LeaseID).
112
Str("name", registerReq.Name).
129
- Str("address", resolvedAddr).
130
- Str("address_advertised", registerReq.Address).
113
Bool("tls_enabled", registerReq.TLSEnabled).
114
Msg("[Registry] Lease registered")
115
@@ -249,7 +231,7 @@ func (r *SDKRegistry) HandleRenew(w http.ResponseWriter, req *http.Request) {
231
}
232
233
// Re-register route if needed (e.g., router restarted while lease remained active).
252
- if err := r.registerSNIRoute(entry.Lease.ID, entry.Lease.Name, entry.Lease.Address); err != nil {
234
+ if err := r.registerSNIRoute(entry.Lease.ID, entry.Lease.Name); err != nil {
235
log.Warn().
236
Err(err).
237
Str("lease_id", entry.Lease.ID).
@@ -262,7 +244,7 @@ func (r *SDKRegistry) HandleRenew(w http.ResponseWriter, req *http.Request) {
244
})
245
}
246
265
-func (r *SDKRegistry) registerSNIRoute(leaseID, name, address string) error {
247
+func (r *SDKRegistry) registerSNIRoute(leaseID, name string) error {
248
if r.sniRouter == nil {
249
return nil
250
}
@@ -270,39 +252,7 @@ func (r *SDKRegistry) registerSNIRoute(leaseID, name, address string) error {
252
return fmt.Errorf("invalid app domain configuration")
253
}
254
sniName := strings.ToLower(strings.TrimSpace(name)) + "." + r.baseHost
273
- return r.sniRouter.RegisterRoute(sniName, address, leaseID, name)
274
-}
275
-
276
-func resolveLeaseAddress(req *http.Request, advertisedAddr string) (string, error) {
277
- advertisedAddr = strings.TrimSpace(advertisedAddr)
278
- host, port, err := net.SplitHostPort(advertisedAddr)
279
- if err != nil {
280
- return "", fmt.Errorf("invalid address: %q", advertisedAddr)
281
- }
282
-
283
- if isLoopbackOrLocalHost(host) {
284
- clientIP := strings.TrimSpace(manager.ExtractClientIP(req))
285
- if clientIP == "" {
286
- return "", fmt.Errorf("cannot resolve client IP for address: %q", advertisedAddr)
287
- }
288
- host = clientIP
289
- }
290
-
291
- return net.JoinHostPort(host, port), nil
292
-}
293
-
294
-func isLoopbackOrLocalHost(host string) bool {
295
- h := strings.ToLower(strings.Trim(strings.TrimSpace(host), "[]"))
296
- if h == "" || h == "localhost" {
297
- return true
298
- }
299
-
300
- ip := net.ParseIP(h)
301
- if ip == nil {
302
- return false
303
- }
304
-
305
- return ip.IsLoopback() || ip.IsUnspecified()
255
+ return r.sniRouter.RegisterRoute(sniName, leaseID, name)
256
}
257
258
func (r *SDKRegistry) unregisterSNIRoute(leaseID string) {
cmd/relay-server/serve.go
+1
-1
@@ -260,7 +260,7 @@ func proxyToHTTP(w http.ResponseWriter, r *http.Request, serv *portal.RelayServe
260
log.Error().
261
Err(err).
262
Str("lease", leaseName).
263
- Str("target", entry.Lease.Address).
263
+ Str("lease_id", entry.Lease.ID).
264
Msg("[proxy] failed to connect to backend")
265
http.Error(w, "service unavailable", http.StatusServiceUnavailable)
266
return
portal/lease.go
-1
@@ -21,7 +21,6 @@ type ParsedMetadata struct {
21
type Lease struct {
22
ID string `json:"id"`
23
Name string `json:"name"`
24
- Address string `json:"address"` // Tunnel client address for TCP connection
24
Metadata Metadata `json:"metadata"`
25
Expires time.Time `json:"expires"`
26
TLSEnabled bool `json:"tls_enabled"` // Whether the tunnel client handles TLS termination
portal/lease_test.go
-3
@@ -17,7 +17,6 @@ func TestLeaseManagerDeleteLeaseInvokesCallback(t *testing.T) {
17
lease := &Lease{
18
ID: "lease-1",
19
Name: "app-1",
20
- Address: "127.0.0.1:10001",
20
Expires: time.Now().Add(30 * time.Second),
21
}
22
if !lm.UpdateLease(lease) {
@@ -44,7 +43,6 @@ func TestLeaseManagerCleanupExpiredLeasesInvokesCallback(t *testing.T) {
43
Lease: &Lease{
44
ID: "expired-1",
45
Name: "expired",
47
- Address: "127.0.0.1:10002",
46
Expires: time.Now().Add(-1 * time.Second),
47
},
48
Expires: time.Now().Add(-1 * time.Second),
@@ -53,7 +51,6 @@ func TestLeaseManagerCleanupExpiredLeasesInvokesCallback(t *testing.T) {
51
Lease: &Lease{
52
ID: "active-1",
53
Name: "active",
56
- Address: "127.0.0.1:10003",
54
Expires: time.Now().Add(30 * time.Second),
55
},
56
Expires: time.Now().Add(30 * time.Second),
portal/reverse_hub.go
+22
@@ -47,12 +47,14 @@ func (c *ReverseConn) Wait() {
47
type ReverseHub struct {
48
mu sync.RWMutex
49
pending map[string]chan *ReverseConn
50
+ dropped map[string]struct{} // leases that have been dropped and should reject offers
51
authorizer func(string, string) bool
52
}
53
54
func NewReverseHub() *ReverseHub {
55
return &ReverseHub{
56
pending: make(map[string]chan *ReverseConn),
57
+ dropped: make(map[string]struct{}),
58
}
59
}
60
@@ -60,6 +62,11 @@ func (h *ReverseHub) getOrCreate(leaseID string) chan *ReverseConn {
62
h.mu.Lock()
63
defer h.mu.Unlock()
64
65
+ // Don't create channels for dropped leases
66
+ if _, dropped := h.dropped[leaseID]; dropped {
67
+ return nil
68
+ }
69
+
70
ch, ok := h.pending[leaseID]
71
if ok {
72
return ch
@@ -94,6 +101,10 @@ func (h *ReverseHub) isAuthorized(leaseID, token string) bool {
101
102
func (h *ReverseHub) Offer(leaseID string, conn *ReverseConn) bool {
103
ch := h.getOrCreate(leaseID)
104
+ if ch == nil {
105
+ // Lease was dropped, reject the offer
106
+ return false
107
+ }
108
select {
109
case ch <- conn:
110
return true
@@ -164,12 +175,15 @@ func (h *ReverseHub) DropLease(leaseID string) {
175
if ok {
176
delete(h.pending, leaseID)
177
}
178
+ // Mark as dropped to prevent new offers from creating channels
179
+ h.dropped[leaseID] = struct{}{}
180
h.mu.Unlock()
181
182
if !ok {
183
return
184
}
185
186
+ // Drain and close any pending connections
187
for {
188
select {
189
case conn := <-ch:
@@ -182,6 +196,14 @@ func (h *ReverseHub) DropLease(leaseID string) {
196
}
197
}
198
199
+// ClearDropped removes a lease from the dropped set, allowing it to be re-registered.
200
+// This should be called when a lease is re-registered after being dropped.
201
+func (h *ReverseHub) ClearDropped(leaseID string) {
202
+ h.mu.Lock()
203
+ delete(h.dropped, leaseID)
204
+ h.mu.Unlock()
205
+}
206
+
207
func (h *ReverseHub) HandleConnect(ws *websocket.Conn) {
208
ws.PayloadType = websocket.BinaryFrame
209
portal/utils/sni/parser.go
+50
-1
@@ -17,6 +17,8 @@ var (
17
ErrNotClientHello = errors.New("not a ClientHello message")
18
// ErrNoSNI is returned when the ClientHello doesn't contain SNI
19
ErrNoSNI = errors.New("no SNI found in ClientHello")
20
+ // ErrInvalidSNI is returned when the SNI hostname is invalid
21
+ ErrInvalidSNI = errors.New("invalid SNI hostname")
22
)
23
24
// ExtractSNI extracts the SNI hostname from a TLS ClientHello message.
@@ -207,7 +209,11 @@ func parseSNIExtension(data []byte) (string, error) {
209
if nameLen == 0 {
210
return "", ErrNoSNI
211
}
210
- return string(data[offset : offset+nameLen]), nil
212
+ hostname := string(data[offset : offset+nameLen])
213
+ if !isValidSNIHostname(hostname) {
214
+ return "", ErrInvalidSNI
215
+ }
216
+ return hostname, nil
217
}
218
219
offset += nameLen
@@ -216,6 +222,49 @@ func parseSNIExtension(data []byte) (string, error) {
222
return "", ErrNoSNI
223
}
224
225
+// isValidSNIHostname validates that a hostname is a valid DNS name per RFC 1035 and RFC 1123.
226
+// - Total length must not exceed 253 characters
227
+// - Labels must be 1-63 characters
228
+// - Labels can contain a-z, A-Z, 0-9, and hyphen
229
+// - Labels cannot start or end with hyphen
230
+// - No null bytes or other control characters
231
+func isValidSNIHostname(hostname string) bool {
232
+ if len(hostname) == 0 || len(hostname) > 253 {
233
+ return false
234
+ }
235
+
236
+ // Check for null bytes and other control characters
237
+ for i := 0; i < len(hostname); i++ {
238
+ if hostname[i] < 0x20 || hostname[i] > 0x7E {
239
+ return false
240
+ }
241
+ }
242
+
243
+ start := 0
244
+ for i := 0; i <= len(hostname); i++ {
245
+ if i == len(hostname) || hostname[i] == '.' {
246
+ label := hostname[start:i]
247
+ if len(label) == 0 || len(label) > 63 {
248
+ return false
249
+ }
250
+ // Check label characters
251
+ for j, c := range []byte(label) {
252
+ // Allow a-z, A-Z, 0-9, and hyphen
253
+ if !((c >= 'a' && c <= 'z') || (c >= 'A' && c <= 'Z') || (c >= '0' && c <= '9') || c == '-') {
254
+ return false
255
+ }
256
+ // Label cannot start or end with hyphen
257
+ if c == '-' && (j == 0 || j == len(label)-1) {
258
+ return false
259
+ }
260
+ }
261
+ start = i + 1
262
+ }
263
+ }
264
+
265
+ return true
266
+}
267
+
268
// PeekSNI peeks at the SNI from a connection without consuming the data.
269
// It returns the SNI and a new reader that includes the peeked data.
270
// This is useful for routing connections before fully reading them.
portal/utils/sni/router.go
+31
-70
@@ -28,10 +28,9 @@ const (
28
29
// Route represents a registered route
30
type Route struct {
31
- SNI string
32
- TargetAddr string
33
- LeaseID string
34
- LeaseName string
31
+ SNI string
32
+ LeaseID string
33
+ LeaseName string
34
}
35
36
// Router handles SNI-based TCP routing
@@ -66,7 +65,7 @@ func (r *Router) SetConnectionCallback(cb func(conn net.Conn, route *Route)) {
65
}
66
67
// RegisterRoute registers a new route for an SNI
69
-func (r *Router) RegisterRoute(sni, targetAddr, leaseID, leaseName string) error {
68
+func (r *Router) RegisterRoute(sni, leaseID, leaseName string) error {
69
r.mu.Lock()
70
defer r.mu.Unlock()
71
@@ -81,20 +80,25 @@ func (r *Router) RegisterRoute(sni, targetAddr, leaseID, leaseName string) error
80
return fmt.Errorf("sni is required")
81
}
82
84
- route := &Route{
85
- SNI: sni,
86
- TargetAddr: targetAddr,
87
- LeaseID: leaseID,
88
- LeaseName: leaseName,
89
- }
90
-
83
// Remove previous SNI entry when a lease is re-registered with a new name.
84
if oldRoute, ok := r.leases[leaseID]; ok && oldRoute.SNI != sni {
85
delete(r.routes, oldRoute.SNI)
86
}
95
- // Keep lease index consistent when SNI is reassigned to another lease.
96
- if oldRoute, ok := r.routes[sni]; ok && oldRoute.LeaseID != leaseID {
97
- delete(r.leases, oldRoute.LeaseID)
87
+
88
+ // Warn if SNI is already registered to a different lease (should not happen if lease manager is consistent)
89
+ if existingRoute, ok := r.routes[sni]; ok && existingRoute.LeaseID != leaseID {
90
+ log.Warn().
91
+ Str("sni", sni).
92
+ Str("existing_lease_id", existingRoute.LeaseID).
93
+ Str("new_lease_id", leaseID).
94
+ Msg("[SNI] SNI already registered to different lease; overwriting")
95
+ delete(r.leases, existingRoute.LeaseID)
96
+ }
97
+
98
+ route := &Route{
99
+ SNI: sni,
100
+ LeaseID: leaseID,
101
+ LeaseName: leaseName,
102
}
103
104
r.routes[sni] = route
@@ -102,7 +106,6 @@ func (r *Router) RegisterRoute(sni, targetAddr, leaseID, leaseName string) error
106
107
log.Info().
108
Str("sni", sni).
105
- Str("target", targetAddr).
109
Str("lease_id", leaseID).
110
Msg("[SNI] Route registered")
111
@@ -153,10 +156,13 @@ func (r *Router) GetRoute(sni string) (*Route, bool) {
156
return route, true
157
}
158
156
- // Try wildcard match (e.g., *.example.com)
159
+ // Try wildcard match (e.g., *.example.com matches foo.example.com)
160
+ // TLS wildcards only match a single DNS label, so for "foo.bar.example.com"
161
+ // we only check "*.bar.example.com", not "*.example.com"
162
parts := strings.Split(sni, ".")
158
- for i := 1; i < len(parts); i++ {
159
- wildcard := "*." + strings.Join(parts[i:], ".")
163
+ if len(parts) >= 2 {
164
+ // Only check the immediate parent wildcard
165
+ wildcard := "*." + strings.Join(parts[1:], ".")
166
if route, ok := r.routes[wildcard]; ok {
167
return route, true
168
}
@@ -290,7 +296,7 @@ func (r *Router) handleConnection(clientConn net.Conn) {
296
297
log.Debug().
298
Str("sni", sni).
293
- Str("target", route.TargetAddr).
299
+ Str("lease_id", route.LeaseID).
300
Str("remote", clientConn.RemoteAddr().String()).
301
Msg("[SNI] Route found")
302
@@ -310,56 +316,11 @@ func (r *Router) handleConnection(clientConn net.Conn) {
316
return
317
}
318
313
- // Default behavior: proxy to target
314
- r.proxyConnection(clientConn, peekedReader, route)
315
-}
316
-
317
-// proxyConnection proxies data between client and target
318
-func (r *Router) proxyConnection(clientConn net.Conn, clientReader io.Reader, route *Route) {
319
- defer clientConn.Close()
320
-
321
- // Connect to target
322
- targetConn, err := net.DialTimeout("tcp", route.TargetAddr, 10*time.Second)
323
- if err != nil {
324
- log.Error().
325
- Err(err).
326
- Str("target", route.TargetAddr).
327
- Str("sni", route.SNI).
328
- Msg("[SNI] Failed to connect to target")
329
- return
330
- }
331
- defer targetConn.Close()
332
-
333
- log.Info().
334
- Str("sni", route.SNI).
335
- Str("target", route.TargetAddr).
336
- Str("client", clientConn.RemoteAddr().String()).
337
- Msg("[SNI] Connection established")
338
-
339
- // Create error channels
340
- errCh := make(chan error, 2)
341
-
342
- // Client -> Target
343
- go func() {
344
- _, err := io.Copy(targetConn, clientReader)
345
- errCh <- err
346
- targetConn.Close()
347
- }()
348
-
349
- // Target -> Client
350
- go func() {
351
- _, err := io.Copy(clientConn, targetConn)
352
- errCh <- err
353
- clientConn.Close()
354
- }()
355
-
356
- // Wait for either direction to close
357
- <-errCh
358
-
359
- log.Debug().
360
- Str("sni", route.SNI).
361
- Str("target", route.TargetAddr).
362
- Msg("[SNI] Connection closed")
319
+ // No callback configured - close connection
320
+ log.Warn().
321
+ Str("sni", sni).
322
+ Msg("[SNI] No connection callback configured, closing connection")
323
+ clientConn.Close()
324
}
325
326
// BridgeConnections bridges two connections
portal/utils/sni/router_test.go
new
+233
@@ -0,0 +1,233 @@
1
+package sni
2
+
3
+import (
4
+ "testing"
5
+)
6
+
7
+func TestRouter_RegisterRoute(t *testing.T) {
8
+ router := NewRouter()
9
+
10
+ // Test basic registration
11
+ err := router.RegisterRoute("example.com", "lease-1", "test")
12
+ if err != nil {
13
+ t.Fatalf("failed to register route: %v", err)
14
+ }
15
+
16
+ // Test duplicate registration by same lease (should succeed - update)
17
+ err = router.RegisterRoute("example.com", "lease-1", "test")
18
+ if err != nil {
19
+ t.Fatalf("failed to update route: %v", err)
20
+ }
21
+
22
+ route, ok := router.GetRoute("example.com")
23
+ if !ok {
24
+ t.Fatal("route not found")
25
+ }
26
+ if route.LeaseID != "lease-1" {
27
+ t.Errorf("expected lease-1, got %s", route.LeaseID)
28
+ }
29
+}
30
+
31
+func TestRouter_UnregisterRoute(t *testing.T) {
32
+ router := NewRouter()
33
+
34
+ err := router.RegisterRoute("example.com", "lease-1", "test")
35
+ if err != nil {
36
+ t.Fatalf("failed to register route: %v", err)
37
+ }
38
+
39
+ router.UnregisterRoute("example.com")
40
+
41
+ _, ok := router.GetRoute("example.com")
42
+ if ok {
43
+ t.Error("expected route to be unregistered")
44
+ }
45
+}
46
+
47
+func TestRouter_UnregisterRouteByLeaseID(t *testing.T) {
48
+ router := NewRouter()
49
+
50
+ err := router.RegisterRoute("example.com", "lease-1", "test")
51
+ if err != nil {
52
+ t.Fatalf("failed to register route: %v", err)
53
+ }
54
+
55
+ router.UnregisterRouteByLeaseID("lease-1")
56
+
57
+ _, ok := router.GetRoute("example.com")
58
+ if ok {
59
+ t.Error("expected route to be unregistered")
60
+ }
61
+}
62
+
63
+func TestRouter_GetRoute_Wildcard(t *testing.T) {
64
+ router := NewRouter()
65
+
66
+ // Register wildcard route
67
+ err := router.RegisterRoute("*.example.com", "lease-1", "test")
68
+ if err != nil {
69
+ t.Fatalf("failed to register wildcard route: %v", err)
70
+ }
71
+
72
+ tests := []struct {
73
+ sni string
74
+ wantOK bool
75
+ wantName string
76
+ }{
77
+ {"foo.example.com", true, "*.example.com"}, // should match
78
+ {"bar.example.com", true, "*.example.com"}, // should match
79
+ {"example.com", false, ""}, // should NOT match (no subdomain)
80
+ {"foo.bar.example.com", false, ""}, // should NOT match (TLS wildcard only matches one level)
81
+ {"other.com", false, ""}, // should NOT match
82
+ }
83
+
84
+ for _, tt := range tests {
85
+ t.Run(tt.sni, func(t *testing.T) {
86
+ route, ok := router.GetRoute(tt.sni)
87
+ if ok != tt.wantOK {
88
+ t.Errorf("GetRoute(%q) = %v, want %v", tt.sni, ok, tt.wantOK)
89
+ return
90
+ }
91
+ if ok && route.SNI != tt.wantName {
92
+ t.Errorf("GetRoute(%q) matched %q, want %q", tt.sni, route.SNI, tt.wantName)
93
+ }
94
+ })
95
+ }
96
+}
97
+
98
+func TestRouter_GetRoute_ExactBeforeWildcard(t *testing.T) {
99
+ router := NewRouter()
100
+
101
+ // Register both exact and wildcard routes
102
+ err := router.RegisterRoute("*.example.com", "lease-1", "wildcard")
103
+ if err != nil {
104
+ t.Fatalf("failed to register wildcard route: %v", err)
105
+ }
106
+
107
+ err = router.RegisterRoute("specific.example.com", "lease-2", "specific")
108
+ if err != nil {
109
+ t.Fatalf("failed to register specific route: %v", err)
110
+ }
111
+
112
+ // Exact match should take precedence
113
+ route, ok := router.GetRoute("specific.example.com")
114
+ if !ok {
115
+ t.Fatal("route not found")
116
+ }
117
+ if route.LeaseID != "lease-2" {
118
+ t.Errorf("expected lease-2 (exact match), got %s", route.LeaseID)
119
+ }
120
+
121
+ // Other subdomains should match wildcard
122
+ route, ok = router.GetRoute("other.example.com")
123
+ if !ok {
124
+ t.Fatal("route not found")
125
+ }
126
+ if route.LeaseID != "lease-1" {
127
+ t.Errorf("expected lease-1 (wildcard match), got %s", route.LeaseID)
128
+ }
129
+}
130
+
131
+func TestRouter_GetRouteByLeaseID(t *testing.T) {
132
+ router := NewRouter()
133
+
134
+ err := router.RegisterRoute("example.com", "lease-1", "test")
135
+ if err != nil {
136
+ t.Fatalf("failed to register route: %v", err)
137
+ }
138
+
139
+ route, ok := router.GetRouteByLeaseID("lease-1")
140
+ if !ok {
141
+ t.Fatal("route not found by lease ID")
142
+ }
143
+ if route.SNI != "example.com" {
144
+ t.Errorf("expected SNI example.com, got %s", route.SNI)
145
+ }
146
+
147
+ _, ok = router.GetRouteByLeaseID("nonexistent")
148
+ if ok {
149
+ t.Error("expected route not found for nonexistent lease ID")
150
+ }
151
+}
152
+
153
+func TestRouter_GetAllRoutes(t *testing.T) {
154
+ router := NewRouter()
155
+
156
+ _ = router.RegisterRoute("example.com", "lease-1", "test1")
157
+ _ = router.RegisterRoute("other.com", "lease-2", "test2")
158
+
159
+ routes := router.GetAllRoutes()
160
+ if len(routes) != 2 {
161
+ t.Errorf("expected 2 routes, got %d", len(routes))
162
+ }
163
+}
164
+
165
+func TestRouter_CaseInsensitive(t *testing.T) {
166
+ router := NewRouter()
167
+
168
+ err := router.RegisterRoute("Example.COM", "lease-1", "test")
169
+ if err != nil {
170
+ t.Fatalf("failed to register route: %v", err)
171
+ }
172
+
173
+ // Should find with different case
174
+ route, ok := router.GetRoute("EXAMPLE.com")
175
+ if !ok {
176
+ t.Fatal("route not found with different case")
177
+ }
178
+ if route.SNI != "example.com" {
179
+ t.Errorf("expected normalized SNI example.com, got %s", route.SNI)
180
+ }
181
+}
182
+
183
+func TestRouter_LeaseRename(t *testing.T) {
184
+ router := NewRouter()
185
+
186
+ // Register with name1
187
+ err := router.RegisterRoute("name1.example.com", "lease-1", "name1")
188
+ if err != nil {
189
+ t.Fatalf("failed to register route: %v", err)
190
+ }
191
+
192
+ // Same lease re-registers with name2
193
+ err = router.RegisterRoute("name2.example.com", "lease-1", "name2")
194
+ if err != nil {
195
+ t.Fatalf("failed to re-register route: %v", err)
196
+ }
197
+
198
+ // Old name should be gone
199
+ _, ok := router.GetRoute("name1.example.com")
200
+ if ok {
201
+ t.Error("old route should be removed")
202
+ }
203
+
204
+ // New name should exist
205
+ route, ok := router.GetRoute("name2.example.com")
206
+ if !ok {
207
+ t.Fatal("new route not found")
208
+ }
209
+ if route.LeaseID != "lease-1" {
210
+ t.Errorf("expected lease-1, got %s", route.LeaseID)
211
+ }
212
+}
213
+
214
+func TestRouter_Stop(t *testing.T) {
215
+ router := NewRouter()
216
+
217
+ err := router.RegisterRoute("example.com", "lease-1", "test")
218
+ if err != nil {
219
+ t.Fatalf("failed to register route: %v", err)
220
+ }
221
+
222
+ // Stop should not panic
223
+ err = router.Stop()
224
+ if err != nil {
225
+ t.Errorf("stop failed: %v", err)
226
+ }
227
+
228
+ // Registration after stop should fail
229
+ err = router.RegisterRoute("other.com", "lease-2", "test2")
230
+ if err != ErrRouterClosed {
231
+ t.Errorf("expected ErrRouterClosed, got %v", err)
232
+ }
233
+}
sdk/client.go
+1
-2
@@ -32,7 +32,7 @@ type Client struct {
32
func NewClient(opt ...ClientOption) (*Client, error) {
33
config := &ClientConfig{
34
BootstrapServers: []string{},
35
- ReverseWorkers: 2,
35
+ ReverseWorkers: 0, // uses defaultReverseWorkers from listener
36
ReverseDialTimeout: 5 * time.Second,
37
}
38
@@ -81,7 +81,6 @@ func (c *Client) Listen(name string, options ...MetadataOption) (net.Listener, e
81
lease := &portal.Lease{
82
ID: generateID(),
83
Name: name,
84
- Address: "",
84
TLSEnabled: c.config.TLSEnabled,
85
ReverseToken: reverseToken,
86
Metadata: portal.Metadata{
sdk/listener.go
+10
-91
@@ -24,9 +24,7 @@ import (
24
const (
25
relayKeepaliveInterval = 10 * time.Second
26
reverseReadTimeout = 1 * time.Second
27
- reverseStartMarker = byte(0x01)
28
- defaultReverseWorkers = 2
29
- maxReverseWorkers = 16
27
+ defaultReverseWorkers = 16
28
defaultReverseDialTimeout = 5 * time.Second
29
)
30
@@ -39,7 +37,6 @@ type Listener struct {
37
httpClient *http.Client
38
39
mu sync.RWMutex
42
- listener net.Listener
40
closed bool
41
acceptCh chan net.Conn
42
reverseWorkers int
@@ -80,9 +77,6 @@ func NewListener(relayAddr string, lease *portal.Lease, tlsConfig *tls.Config, a
77
if reverseWorkers <= 0 {
78
reverseWorkers = defaultReverseWorkers
79
}
83
- if reverseWorkers > maxReverseWorkers {
84
- reverseWorkers = maxReverseWorkers
85
- }
80
if reverseDialTimeout <= 0 {
81
reverseDialTimeout = defaultReverseDialTimeout
82
}
@@ -102,51 +96,21 @@ func NewListener(relayAddr string, lease *portal.Lease, tlsConfig *tls.Config, a
96
}, nil
97
}
98
105
-// Start initializes the local listener, then registers it to relay.
99
+// Start registers the lease with relay and starts reverse workers.
100
func (l *Listener) Start() error {
101
l.mu.Lock()
102
if l.closed {
103
l.mu.Unlock()
104
return net.ErrClosed
105
}
112
- if l.listener != nil {
113
- l.mu.Unlock()
114
- return nil
115
- }
106
l.mu.Unlock()
107
118
- rawListener, err := net.Listen("tcp", ":0")
119
- if err != nil {
120
- return fmt.Errorf("listen local tunnel socket: %w", err)
121
- }
122
-
123
- addr, ok := rawListener.Addr().(*net.TCPAddr)
124
- if !ok {
125
- rawListener.Close()
126
- return fmt.Errorf("unexpected listener address type: %T", rawListener.Addr())
127
- }
128
-
129
- // Use localhost for local development. For production/Docker, this should be configurable.
130
- tunnelAddr := fmt.Sprintf("127.0.0.1:%d", addr.Port)
131
- if err := l.registerWithRelay(tunnelAddr); err != nil {
132
- rawListener.Close()
108
+ if err := l.registerWithRelay(); err != nil {
109
return fmt.Errorf("register lease with relay: %w", err)
110
}
111
136
- l.mu.Lock()
137
- if l.closed {
138
- l.mu.Unlock()
139
- rawListener.Close()
140
- return net.ErrClosed
141
- }
142
- l.lease.Address = tunnelAddr
143
- l.listener = rawListener
144
- l.mu.Unlock()
145
-
112
l.wg.Add(1)
113
go l.keepaliveLoop()
148
- l.wg.Add(1)
149
- go l.localAcceptLoop(rawListener)
114
for i := 0; i < l.reverseWorkers; i++ {
115
l.wg.Add(1)
116
go l.reverseAcceptWorker(i)
@@ -156,7 +120,6 @@ func (l *Listener) Start() error {
120
Str("lease_id", l.lease.ID).
121
Str("name", l.lease.Name).
122
Str("relay", l.relayAddr).
159
- Str("tunnel_addr", tunnelAddr).
123
Int("reverse_workers", l.reverseWorkers).
124
Msg("[SDK] Relay listener started")
125
@@ -197,7 +160,7 @@ func (l *Listener) Accept() (net.Conn, error) {
160
return conn, nil
161
}
162
200
-// Close unregisters lease from relay and closes local listener.
163
+// Close unregisters lease from relay.
164
func (l *Listener) Close() error {
165
var retErr error
166
l.closeOnce.Do(func() {
@@ -205,35 +168,22 @@ func (l *Listener) Close() error {
168
169
l.mu.Lock()
170
l.closed = true
208
- listener := l.listener
209
- l.listener = nil
171
l.mu.Unlock()
172
212
- if listener != nil {
213
- retErr = listener.Close()
214
- }
215
-
173
l.wg.Wait()
174
175
if err := l.unregisterFromRelay(); err != nil {
176
log.Warn().Err(err).Str("lease_id", l.lease.ID).Msg("[SDK] Failed to unregister lease")
220
- if retErr == nil {
221
- retErr = err
222
- }
177
+ retErr = err
178
}
179
})
180
181
return retErr
182
}
183
229
-// Addr returns local listener address.
184
+// Addr returns a dummy address (connections come via reverse tunnel).
185
func (l *Listener) Addr() net.Addr {
231
- l.mu.RLock()
232
- defer l.mu.RUnlock()
233
- if l.listener != nil {
234
- return l.listener.Addr()
235
- }
236
- return nil
186
+ return &net.TCPAddr{IP: net.IPv4(0, 0, 0, 0), Port: 0}
187
}
188
189
// LeaseID returns lease ID registered to relay.
@@ -273,30 +223,6 @@ func (l *Listener) keepaliveLoop() {
223
}
224
}
225
276
-func (l *Listener) localAcceptLoop(listener net.Listener) {
277
- defer l.wg.Done()
278
-
279
- for {
280
- conn, err := listener.Accept()
281
- if err != nil {
282
- select {
283
- case <-l.stopCh:
284
- return
285
- default:
286
- log.Warn().Err(err).Str("lease_id", l.lease.ID).Msg("[SDK] Local accept error")
287
- continue
288
- }
289
- }
290
-
291
- select {
292
- case <-l.stopCh:
293
- conn.Close()
294
- return
295
- case l.acceptCh <- conn:
296
- }
297
- }
298
-}
299
-
226
func (l *Listener) reverseAcceptWorker(workerID int) {
227
defer l.wg.Done()
228
@@ -370,7 +296,7 @@ func (l *Listener) waitForReverseStart(conn net.Conn) error {
296
_, err := io.ReadFull(conn, marker[:])
297
if err == nil {
298
_ = conn.SetReadDeadline(time.Time{})
373
- if marker[0] != reverseStartMarker {
299
+ if marker[0] != portal.ReverseStartMarker {
300
return fmt.Errorf("invalid reverse marker: %d", marker[0])
301
}
302
return nil
@@ -394,11 +320,10 @@ func (l *Listener) waitForReverseStart(conn net.Conn) error {
320
}
321
}
322
397
-func (l *Listener) registerWithRelay(tunnelAddr string) error {
323
+func (l *Listener) registerWithRelay() error {
324
reqBody := RegisterRequest{
325
LeaseID: l.lease.ID,
326
Name: l.lease.Name,
401
- Address: tunnelAddr,
327
Metadata: l.lease.Metadata,
328
TLSEnabled: l.lease.TLSEnabled,
329
ReverseToken: l.lease.ReverseToken,
@@ -423,13 +348,7 @@ func (l *Listener) sendKeepalive() error {
348
}
349
350
func (l *Listener) reRegisterLease() error {
426
- l.mu.RLock()
427
- addr := strings.TrimSpace(l.lease.Address)
428
- l.mu.RUnlock()
429
- if addr == "" {
430
- return fmt.Errorf("lease address is empty; cannot re-register")
431
- }
432
- return l.registerWithRelay(addr)
351
+ return l.registerWithRelay()
352
}
353
354
func (l *Listener) postJSON(path string, body any) error {
sdk/types.go
+1
-2
@@ -26,7 +26,7 @@ type ClientConfig struct {
26
HealthCheckInterval time.Duration // Interval for health checks (default: 10 seconds)
27
ReconnectMaxRetries int // Maximum reconnection attempts (default: 0 = infinite)
28
ReconnectInterval time.Duration // Interval between reconnection attempts (default: 5 seconds)
29
- ReverseWorkers int // Number of reverse websocket workers per listener (default: 2)
29
+ ReverseWorkers int // Number of reverse websocket workers per listener (default: 16)
30
ReverseDialTimeout time.Duration // Reverse websocket dial timeout (default: 5 seconds)
31
32
// TLS configuration for tunnel server mode
@@ -150,7 +150,6 @@ func WithHide(hide bool) MetadataOption {
150
type RegisterRequest struct {
151
LeaseID string `json:"lease_id"`
152
Name string `json:"name"`
153
- Address string `json:"address"` // Backend address for TCP connection
153
Metadata portal.Metadata `json:"metadata"`
154
TLSEnabled bool `json:"tls_enabled"` // Whether the backend handles TLS termination
155
ReverseToken string `json:"reverse_token"`