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"`