refactor(backend): consolidate relay routing and connection guards
cognitive committed
Mar 3, 2026 at 20:03 UTC
095ce7723fc4e44aa5f1155d05279225b71c2f6d
10 files changed
+435
-92
.golangci.yml
+3
@@ -43,6 +43,9 @@ linters:
43
- mirror
44
45
exclusions:
46
+ paths:
47
+ - ^cmd/relay-server/frontend/node_modules/
48
+ - ^cmd/relay-server/dist/
49
rules:
50
- linters: [errcheck]
51
source: "^\\s*defer\\s+"
cmd/relay-server/main.go
+39
-27
@@ -147,44 +147,35 @@ func runServer(cfg relayServerConfig) error {
147
148
// Set up SNI connection callback to route to tunnel backends
149
serv.GetSNIRouter().SetConnectionCallback(func(clientConn net.Conn, route *sni.Route) {
150
- if _, ok := serv.GetLeaseManager().GetLeaseByID(route.LeaseID); !ok {
151
- log.Warn().
152
- Str("lease_id", route.LeaseID).
153
- Str("sni", route.SNI).
154
- Msg("[SNI] Lease not active; dropping connection and unregistering route")
155
- serv.GetSNIRouter().UnregisterRouteByLeaseID(route.LeaseID)
156
- if err := clientConn.Close(); err != nil {
157
- log.Debug().
158
- Err(err).
159
- Str("lease_id", route.LeaseID).
160
- Str("sni", route.SNI).
161
- Msg("[SNI] failed to close client connection")
162
- }
150
+ leaseID := ""
151
+ if route != nil {
152
+ leaseID = strings.TrimSpace(route.LeaseID)
153
+ }
154
+ if leaseID == "" {
155
+ logSNIRouteWarning(route, nil, "[SNI] Missing lease id in route; dropping connection")
156
+ closeSNIClientConn(clientConn, route)
157
+ return
158
+ }
159
+
160
+ if _, ok := serv.GetLeaseManager().GetLeaseByID(leaseID); !ok {
161
+ logSNIRouteWarning(route, nil, "[SNI] Lease not active; dropping connection and unregistering route")
162
+ serv.GetSNIRouter().UnregisterRouteByLeaseID(leaseID)
163
+ closeSNIClientConn(clientConn, route)
164
return
165
}
166
167
// Get BPS manager for rate limiting
168
bpsManager := admin.GetBPSManager()
169
169
- reverseConn, err := serv.GetReverseHub().AcquireForTLS(route.LeaseID, portal.TLSAcquireWait)
170
+ reverseConn, err := serv.GetReverseHub().AcquireForTLS(leaseID, portal.TLSAcquireWait)
171
if err != nil {
171
- log.Warn().
172
- Err(err).
173
- Str("lease_id", route.LeaseID).
174
- Str("sni", route.SNI).
175
- Msg("[SNI] Reverse tunnel unavailable")
176
- if err := clientConn.Close(); err != nil {
177
- log.Debug().
178
- Err(err).
179
- Str("lease_id", route.LeaseID).
180
- Str("sni", route.SNI).
181
- Msg("[SNI] failed to close client connection")
182
- }
172
+ logSNIRouteWarning(route, err, "[SNI] Reverse tunnel unavailable")
173
+ closeSNIClientConn(clientConn, route)
174
return
175
}
176
177
// SNI path is reverse-only (NAT-friendly): relay never dials app directly.
187
- manager.EstablishRelayWithBPS(clientConn, reverseConn.Conn, route.LeaseID, bpsManager)
178
+ manager.EstablishRelayWithBPS(clientConn, reverseConn.Conn, leaseID, bpsManager)
179
reverseConn.Close()
180
})
181
@@ -212,6 +203,27 @@ func runServer(cfg relayServerConfig) error {
203
return nil
204
}
205
206
+func closeSNIClientConn(clientConn net.Conn, route *sni.Route) {
207
+ if err := clientConn.Close(); err != nil {
208
+ withSNIRouteFields(log.Debug().Err(err), route).Msg("[SNI] failed to close client connection")
209
+ }
210
+}
211
+
212
+func logSNIRouteWarning(route *sni.Route, err error, msg string) {
213
+ event := log.Warn()
214
+ if err != nil {
215
+ event = event.Err(err)
216
+ }
217
+ withSNIRouteFields(event, route).Msg(msg)
218
+}
219
+
220
+func withSNIRouteFields(event *zerolog.Event, route *sni.Route) *zerolog.Event {
221
+ if route == nil {
222
+ return event
223
+ }
224
+ return event.Str("lease_id", route.LeaseID).Str("sni", route.SNI)
225
+}
226
+
227
func trimmedEnv(name string) string {
228
return strings.TrimSpace(os.Getenv(name))
229
}
cmd/relay-server/registry.go
+61
-33
@@ -32,11 +32,19 @@ func reverseTokenMatches(expected, provided string) bool {
32
return subtle.ConstantTimeCompare([]byte(expected), []byte(provided)) == 1
33
}
34
35
+func normalizeLeaseID(raw string) string {
36
+ return strings.TrimSpace(raw)
37
+}
38
+
39
+func normalizeLeaseCredentials(leaseID, reverseToken string) (string, string) {
40
+ return normalizeLeaseID(leaseID), strings.TrimSpace(reverseToken)
41
+}
42
+
43
func lookupLeaseEntry(serv *portal.RelayServer, leaseID string) (*portal.LeaseEntry, bool) {
44
if serv == nil {
45
return nil, false
46
}
39
- entry, ok := serv.GetLeaseManager().GetLeaseByID(strings.TrimSpace(leaseID))
47
+ entry, ok := serv.GetLeaseManager().GetLeaseByID(normalizeLeaseID(leaseID))
48
if !ok || entry == nil || entry.Lease == nil {
49
return nil, false
50
}
@@ -64,6 +72,34 @@ func (r *SDKRegistry) requireMethod(w http.ResponseWriter, req *http.Request, me
72
return false
73
}
74
75
+func (r *SDKRegistry) decodeRequestBody(w http.ResponseWriter, req *http.Request, dst any, logMessage string) bool {
76
+ if err := json.NewDecoder(req.Body).Decode(dst); err != nil {
77
+ log.Error().Err(err).Msg(logMessage)
78
+ writeAPIError(w, http.StatusBadRequest, "invalid_request", "invalid request body")
79
+ return false
80
+ }
81
+ return true
82
+}
83
+
84
+func (r *SDKRegistry) validateLeaseCredentials(w http.ResponseWriter, leaseID, reverseToken string) bool {
85
+ if leaseID == "" {
86
+ writeAPIError(w, http.StatusBadRequest, "missing_lease_id", "lease_id is required")
87
+ return false
88
+ }
89
+ if reverseToken == "" {
90
+ writeAPIError(w, http.StatusBadRequest, "missing_reverse_token", "reverse_token is required")
91
+ return false
92
+ }
93
+ return true
94
+}
95
+
96
+func isWebSocketUpgrade(req *http.Request) bool {
97
+ if req == nil {
98
+ return false
99
+ }
100
+ return strings.EqualFold(strings.TrimSpace(req.Header.Get("Upgrade")), "websocket")
101
+}
102
+
103
// HandleSDKRequest routes /sdk/* requests.
104
func (r *SDKRegistry) HandleSDKRequest(w http.ResponseWriter, req *http.Request, serv *portal.RelayServer) {
105
path := strings.TrimSuffix(req.URL.Path, "/")
@@ -89,17 +125,24 @@ func (r *SDKRegistry) handleConnect(w http.ResponseWriter, req *http.Request, se
125
return
126
}
127
92
- leaseID := strings.TrimSpace(req.URL.Query().Get("lease_id"))
128
+ leaseID, token := normalizeLeaseCredentials(
129
+ req.URL.Query().Get("lease_id"),
130
+ req.Header.Get(portal.ReverseConnectTokenHeader),
131
+ )
132
if leaseID == "" {
133
http.Error(w, "missing lease_id", http.StatusBadRequest)
134
return
135
}
97
- token := strings.TrimSpace(req.Header.Get(portal.ReverseConnectTokenHeader))
136
if token == "" {
137
http.Error(w, "missing reverse token", http.StatusUnauthorized)
138
return
139
}
102
- clientIP := r.extractClientIP(req)
140
+ if isWebSocketUpgrade(req) {
141
+ http.Error(w, "websocket transport is not supported", http.StatusBadRequest)
142
+ return
143
+ }
144
+
145
+ clientIP := strings.TrimSpace(r.extractClientIP(req))
146
if r.isClientIPBanned(clientIP) {
147
http.Error(w, "ip is banned", http.StatusForbidden)
148
return
@@ -126,11 +169,15 @@ func (r *SDKRegistry) handleConnect(w http.ResponseWriter, req *http.Request, se
169
return
170
}
171
if _, err := rw.WriteString("HTTP/1.1 200 OK\r\nContent-Length: 0\r\nConnection: keep-alive\r\n\r\n"); err != nil {
129
- _ = conn.Close()
172
+ if closeErr := conn.Close(); closeErr != nil {
173
+ log.Debug().Err(closeErr).Msg("[Registry] failed to close hijacked connection after write failure")
174
+ }
175
return
176
}
177
if err := rw.Flush(); err != nil {
133
- _ = conn.Close()
178
+ if closeErr := conn.Close(); closeErr != nil {
179
+ log.Debug().Err(closeErr).Msg("[Registry] failed to close hijacked connection after flush failure")
180
+ }
181
return
182
}
183
@@ -144,23 +191,14 @@ func (r *SDKRegistry) handleRegister(w http.ResponseWriter, req *http.Request, s
191
}
192
193
var registerReq types.RegisterRequest
147
- if err := json.NewDecoder(req.Body).Decode(®isterReq); err != nil {
148
- log.Error().Err(err).Msg("[Registry] Failed to decode registration request")
149
- writeAPIError(w, http.StatusBadRequest, "invalid_request", "invalid request body")
194
+ if !r.decodeRequestBody(w, req, ®isterReq, "[Registry] Failed to decode registration request") {
195
return
196
}
197
153
- registerReq.LeaseID = strings.TrimSpace(registerReq.LeaseID)
198
+ registerReq.LeaseID, registerReq.ReverseToken = normalizeLeaseCredentials(registerReq.LeaseID, registerReq.ReverseToken)
199
registerReq.Name = strings.TrimSpace(registerReq.Name)
155
- registerReq.ReverseToken = strings.TrimSpace(registerReq.ReverseToken)
200
157
- if registerReq.LeaseID == "" {
158
- writeAPIError(w, http.StatusBadRequest, "missing_lease_id", "lease_id is required")
159
- return
160
- }
161
-
162
- if registerReq.ReverseToken == "" {
163
- writeAPIError(w, http.StatusBadRequest, "missing_reverse_token", "reverse_token is required")
201
+ if !r.validateLeaseCredentials(w, registerReq.LeaseID, registerReq.ReverseToken) {
202
return
203
}
204
name := registerReq.Name
@@ -240,12 +278,10 @@ func (r *SDKRegistry) handleUnregister(w http.ResponseWriter, req *http.Request,
278
}
279
280
var unregisterReq types.UnregisterRequest
243
- if err := json.NewDecoder(req.Body).Decode(&unregisterReq); err != nil {
244
- log.Error().Err(err).Msg("[Registry] Failed to decode unregistration request")
245
- writeAPIError(w, http.StatusBadRequest, "invalid_request", "invalid request body")
281
+ if !r.decodeRequestBody(w, req, &unregisterReq, "[Registry] Failed to decode unregistration request") {
282
return
283
}
248
- unregisterReq.LeaseID = strings.TrimSpace(unregisterReq.LeaseID)
284
+ unregisterReq.LeaseID = normalizeLeaseID(unregisterReq.LeaseID)
285
286
// Delete from lease manager
287
if serv.GetLeaseManager().DeleteLease(unregisterReq.LeaseID) {
@@ -266,20 +302,12 @@ func (r *SDKRegistry) handleRenew(w http.ResponseWriter, req *http.Request, serv
302
}
303
304
var renewReq types.RenewRequest
269
- if err := json.NewDecoder(req.Body).Decode(&renewReq); err != nil {
270
- log.Error().Err(err).Msg("[Registry] Failed to decode renewal request")
271
- writeAPIError(w, http.StatusBadRequest, "invalid_request", "invalid request body")
305
+ if !r.decodeRequestBody(w, req, &renewReq, "[Registry] Failed to decode renewal request") {
306
return
307
}
308
275
- renewReq.LeaseID = strings.TrimSpace(renewReq.LeaseID)
276
- renewReq.ReverseToken = strings.TrimSpace(renewReq.ReverseToken)
277
- if renewReq.LeaseID == "" {
278
- writeAPIError(w, http.StatusBadRequest, "missing_lease_id", "lease_id is required")
279
- return
280
- }
281
- if renewReq.ReverseToken == "" {
282
- writeAPIError(w, http.StatusBadRequest, "missing_reverse_token", "reverse_token is required")
309
+ renewReq.LeaseID, renewReq.ReverseToken = normalizeLeaseCredentials(renewReq.LeaseID, renewReq.ReverseToken)
310
+ if !r.validateLeaseCredentials(w, renewReq.LeaseID, renewReq.ReverseToken) {
311
return
312
}
313
cmd/relay-server/registry_test.go
+48
@@ -135,3 +135,51 @@ func TestSDKRegistryHandleConnectRejectsBannedIP(t *testing.T) {
135
t.Fatalf("expected banned ip error body, got %q", rec.Body.String())
136
}
137
}
138
+
139
+func TestIsWebSocketUpgrade(t *testing.T) {
140
+ t.Parallel()
141
+
142
+ tests := []struct {
143
+ name string
144
+ header string
145
+ want bool
146
+ }{
147
+ {name: "empty", header: "", want: false},
148
+ {name: "websocket lowercase", header: "websocket", want: true},
149
+ {name: "websocket mixed case", header: "WebSocket", want: true},
150
+ {name: "other upgrade", header: "h2c", want: false},
151
+ }
152
+
153
+ for _, tt := range tests {
154
+ t.Run(tt.name, func(t *testing.T) {
155
+ t.Parallel()
156
+
157
+ req := httptest.NewRequest(http.MethodGet, types.PathSDKConnect, http.NoBody)
158
+ if tt.header != "" {
159
+ req.Header.Set("Upgrade", tt.header)
160
+ }
161
+ if got := isWebSocketUpgrade(req); got != tt.want {
162
+ t.Fatalf("isWebSocketUpgrade(%q)=%v, want %v", tt.header, got, tt.want)
163
+ }
164
+ })
165
+ }
166
+}
167
+
168
+func TestSDKRegistryHandleConnectRejectsWebSocketUpgrade(t *testing.T) {
169
+ serv := newRegistryTestRelayServer(t)
170
+ registry := &SDKRegistry{}
171
+
172
+ req := httptest.NewRequest(http.MethodGet, types.PathSDKConnect+"?lease_id=lease-websocket", http.NoBody)
173
+ req.Header.Set(portal.ReverseConnectTokenHeader, "reverse-token")
174
+ req.Header.Set("Upgrade", "websocket")
175
+ rec := httptest.NewRecorder()
176
+
177
+ registry.handleConnect(rec, req, serv)
178
+
179
+ if rec.Code != http.StatusBadRequest {
180
+ t.Fatalf("handleConnect status = %d, want %d", rec.Code, http.StatusBadRequest)
181
+ }
182
+ if !strings.Contains(strings.ToLower(rec.Body.String()), "websocket") {
183
+ t.Fatalf("expected websocket rejection body, got %q", rec.Body.String())
184
+ }
185
+}
portal/reverse_hub.go
+25
@@ -3,6 +3,7 @@ package portal
3
import (
4
"fmt"
5
"net"
6
+ "strings"
7
"sync"
8
"sync/atomic"
9
"time"
@@ -201,6 +202,11 @@ func (h *ReverseHub) notifyAccepted(leaseID, ip string) {
202
}
203
204
func (h *ReverseHub) Offer(leaseID string, conn *ReverseConn) bool {
205
+ leaseID = strings.TrimSpace(leaseID)
206
+ if leaseID == "" || conn == nil || conn.Conn == nil {
207
+ return false
208
+ }
209
+
210
pool := h.getOrCreatePool(leaseID)
211
if pool == nil {
212
return false
@@ -227,6 +233,8 @@ func (h *ReverseHub) Offer(leaseID string, conn *ReverseConn) bool {
233
}
234
235
func (h *ReverseHub) AcquireForTLS(leaseID string, timeout time.Duration) (*ReverseConn, error) {
236
+ leaseID = strings.TrimSpace(leaseID)
237
+
238
pool, ok := h.getPool(leaseID)
239
if !ok {
240
return nil, fmt.Errorf("no tunnel available for lease %s", leaseID)
@@ -278,6 +286,11 @@ func (h *ReverseHub) AcquireForTLS(leaseID string, timeout time.Duration) (*Reve
286
}
287
288
func (h *ReverseHub) DropLease(leaseID string) {
289
+ leaseID = strings.TrimSpace(leaseID)
290
+ if leaseID == "" {
291
+ return
292
+ }
293
+
294
h.mu.Lock()
295
pool, ok := h.pools[leaseID]
296
if ok {
@@ -306,6 +319,11 @@ func (h *ReverseHub) DropLease(leaseID string) {
319
// ClearDropped removes a lease from the dropped set, allowing it to be re-registered.
320
// This should be called when a lease is re-registered after being dropped.
321
func (h *ReverseHub) ClearDropped(leaseID string) {
322
+ leaseID = strings.TrimSpace(leaseID)
323
+ if leaseID == "" {
324
+ return
325
+ }
326
+
327
h.mu.Lock()
328
delete(h.dropped, leaseID)
329
h.mu.Unlock()
@@ -316,6 +334,10 @@ func (h *ReverseHub) HandleConnect(conn net.Conn, leaseID, token, remoteIP strin
334
return
335
}
336
337
+ leaseID = strings.TrimSpace(leaseID)
338
+ token = strings.TrimSpace(token)
339
+ remoteIP = strings.TrimSpace(remoteIP)
340
+
341
if leaseID == "" {
342
log.Warn().Msg("[ReverseHub] Missing lease_id on reverse connect")
343
h.rejectConn(conn, "[ReverseHub] failed to close unauthorized reverse connection")
@@ -357,6 +379,9 @@ func (h *ReverseHub) rejectConn(conn net.Conn, debugCloseMessage string) {
379
}
380
381
func (h *ReverseHub) closeConn(conn net.Conn, debugCloseMessage string) {
382
+ if conn == nil {
383
+ return
384
+ }
385
if err := conn.Close(); err != nil {
386
log.Debug().Err(err).Msg(debugCloseMessage)
387
}
portal/reverse_hub_test.go
+75
@@ -27,6 +27,81 @@ func TestReverseHubAuthorization(t *testing.T) {
27
}
28
}
29
30
+func TestReverseHubOfferRejectsInvalidInput(t *testing.T) {
31
+ hub := NewReverseHub()
32
+
33
+ if ok := hub.Offer("", nil); ok {
34
+ t.Fatal("expected offer with empty lease and nil connection to fail")
35
+ }
36
+ if ok := hub.Offer(" ", nil); ok {
37
+ t.Fatal("expected offer with whitespace lease and nil connection to fail")
38
+ }
39
+}
40
+
41
+func TestHandleConnectTrimsLeaseIDAndToken(t *testing.T) {
42
+ hub := NewReverseHub()
43
+ leaseID := "lease-connect-trim"
44
+ token := "token-connect-trim"
45
+ hub.SetAuthorizer(func(gotLeaseID, gotToken string) bool {
46
+ return gotLeaseID == leaseID && gotToken == token
47
+ })
48
+
49
+ local, peer := net.Pipe()
50
+ defer func() {
51
+ _ = peer.Close()
52
+ }()
53
+
54
+ done := make(chan struct{})
55
+ go func() {
56
+ hub.HandleConnect(local, " "+leaseID+" ", " "+token+" ", " 127.0.0.1 ")
57
+ close(done)
58
+ }()
59
+
60
+ markerRead := make(chan byte, 1)
61
+ readErr := make(chan error, 1)
62
+ go func() {
63
+ var b [1]byte
64
+ _, err := io.ReadFull(peer, b[:])
65
+ if err != nil {
66
+ readErr <- err
67
+ return
68
+ }
69
+ markerRead <- b[0]
70
+ }()
71
+
72
+ got, err := hub.AcquireForTLS(leaseID, 500*time.Millisecond)
73
+ deadline := time.Now().Add(500 * time.Millisecond)
74
+ for err != nil {
75
+ if time.Now().After(deadline) {
76
+ t.Fatalf("AcquireForTLS failed: %v", err)
77
+ }
78
+ time.Sleep(10 * time.Millisecond)
79
+ got, err = hub.AcquireForTLS(leaseID, 25*time.Millisecond)
80
+ }
81
+ if got == nil {
82
+ t.Fatal("AcquireForTLS returned nil connection")
83
+ }
84
+
85
+ select {
86
+ case err := <-readErr:
87
+ t.Fatalf("failed to read marker: %v", err)
88
+ case marker := <-markerRead:
89
+ if marker != TLSStartMarker {
90
+ t.Fatalf("unexpected marker: %d", marker)
91
+ }
92
+ case <-time.After(500 * time.Millisecond):
93
+ t.Fatal("timed out waiting for start marker")
94
+ }
95
+
96
+ got.Close()
97
+
98
+ select {
99
+ case <-done:
100
+ case <-time.After(500 * time.Millisecond):
101
+ t.Fatal("HandleConnect did not return after connection close")
102
+ }
103
+}
104
+
105
func TestAcquireForTLSSendsStartMarker(t *testing.T) {
106
hub := NewReverseHub()
107
leaseID := "lease-tls-marker"
portal/sni/router.go
+28
-23
@@ -235,9 +235,7 @@ func (r *Router) Stop() error {
235
236
r.mu.Lock()
237
if r.listener != nil {
238
- if err := r.listener.Close(); err != nil {
239
- log.Debug().Err(err).Msg("[SNI] failed to close listener")
240
- }
238
+ closeWithDebugLog(r.listener, "[SNI] failed to close listener")
239
}
240
r.mu.Unlock()
241
})
@@ -295,9 +293,7 @@ func (r *Router) handleConnection(clientConn net.Conn) {
293
Err(err).
294
Str("remote", clientConn.RemoteAddr().String()).
295
Msg("[SNI] Failed to extract SNI")
298
- if closeErr := clientConn.Close(); closeErr != nil {
299
- log.Debug().Err(closeErr).Msg("[SNI] failed to close client connection")
300
- }
296
+ closeWithDebugLog(clientConn, "[SNI] failed to close client connection")
297
return
298
}
299
@@ -315,9 +311,7 @@ func (r *Router) handleConnection(clientConn net.Conn) {
311
// Find the route
312
route, ok := r.GetRoute(sni)
313
if !ok {
318
- r.mu.RLock()
319
- onNoRoute := r.onNoRoute
320
- r.mu.RUnlock()
314
+ onNoRoute := r.getNoRouteHandler()
315
if onNoRoute != nil && onNoRoute(wrappedConn, sni) {
316
return
317
}
@@ -326,9 +320,7 @@ func (r *Router) handleConnection(clientConn net.Conn) {
320
Str("sni", sni).
321
Str("remote", clientConn.RemoteAddr().String()).
322
Msg("[SNI] No route found")
329
- if closeErr := clientConn.Close(); closeErr != nil {
330
- log.Debug().Err(closeErr).Msg("[SNI] failed to close unrouted client connection")
331
- }
323
+ closeWithDebugLog(clientConn, "[SNI] failed to close unrouted client connection")
324
return
325
}
326
@@ -339,9 +331,7 @@ func (r *Router) handleConnection(clientConn net.Conn) {
331
Msg("[SNI] Route found")
332
333
// Call the connection callback if set
342
- r.mu.RLock()
343
- onConnection := r.onConnection
344
- r.mu.RUnlock()
334
+ onConnection := r.getConnectionHandler()
335
336
if onConnection != nil {
337
onConnection(wrappedConn, route)
@@ -352,9 +342,28 @@ func (r *Router) handleConnection(clientConn net.Conn) {
342
log.Warn().
343
Str("sni", sni).
344
Msg("[SNI] No connection callback configured, closing connection")
355
- if err := clientConn.Close(); err != nil {
356
- log.Debug().Err(err).Msg("[SNI] failed to close client connection")
345
+ closeWithDebugLog(clientConn, "[SNI] failed to close client connection")
346
+}
347
+
348
+func closeWithDebugLog(closer io.Closer, msg string) {
349
+ if closer == nil {
350
+ return
351
}
352
+ if err := closer.Close(); err != nil {
353
+ log.Debug().Err(err).Msg(msg)
354
+ }
355
+}
356
+
357
+func (r *Router) getNoRouteHandler() func(conn net.Conn, sni string) bool {
358
+ r.mu.RLock()
359
+ defer r.mu.RUnlock()
360
+ return r.onNoRoute
361
+}
362
+
363
+func (r *Router) getConnectionHandler() func(conn net.Conn, route *Route) {
364
+ r.mu.RLock()
365
+ defer r.mu.RUnlock()
366
+ return r.onConnection
367
}
368
369
// BridgeConnections bridges two connections.
@@ -368,18 +377,14 @@ func BridgeConnections(conn1, conn2 net.Conn) {
377
go func() {
378
_, err := io.Copy(conn2, conn1)
379
errCh <- err
371
- if closeErr := conn2.Close(); closeErr != nil {
372
- log.Debug().Err(closeErr).Msg("[SNI] failed to close bridged connection 2")
373
- }
380
+ closeWithDebugLog(conn2, "[SNI] failed to close bridged connection 2")
381
}()
382
383
// Conn2 -> Conn1
384
go func() {
385
_, err := io.Copy(conn1, conn2)
386
errCh <- err
380
- if closeErr := conn1.Close(); closeErr != nil {
381
- log.Debug().Err(closeErr).Msg("[SNI] failed to close bridged connection 1")
382
- }
387
+ closeWithDebugLog(conn1, "[SNI] failed to close bridged connection 1")
388
}()
389
390
// Wait for either direction to close
portal/sni/router_test.go
+113
@@ -2,7 +2,9 @@ package sni
2
3
import (
4
"errors"
5
+ "net"
6
"testing"
7
+ "time"
8
)
9
10
func TestRouter_RegisterRoute(t *testing.T) {
@@ -232,3 +234,114 @@ func TestRouter_Stop(t *testing.T) {
234
t.Errorf("expected ErrRouterClosed, got %v", err)
235
}
236
}
237
+
238
+func TestRouter_HandleConnectionNoRouteHandlerHandled(t *testing.T) {
239
+ router := NewRouter("")
240
+ noRouteCalls := make(chan string, 1)
241
+ router.SetNoRouteHandler(func(_ net.Conn, sni string) bool {
242
+ select {
243
+ case noRouteCalls <- sni:
244
+ default:
245
+ }
246
+ return true
247
+ })
248
+
249
+ client, server := net.Pipe()
250
+ defer func() {
251
+ _ = client.Close()
252
+ _ = server.Close()
253
+ }()
254
+
255
+ done := make(chan struct{})
256
+ router.wg.Add(1)
257
+ go func() {
258
+ router.handleConnection(server)
259
+ close(done)
260
+ }()
261
+
262
+ if _, err := client.Write(buildClientHello("tenant.example.com", true)); err != nil {
263
+ t.Fatalf("write client hello: %v", err)
264
+ }
265
+
266
+ select {
267
+ case gotSNI := <-noRouteCalls:
268
+ if gotSNI != "tenant.example.com" {
269
+ t.Fatalf("no-route handler sni=%q, want %q", gotSNI, "tenant.example.com")
270
+ }
271
+ case <-time.After(500 * time.Millisecond):
272
+ t.Fatal("no-route handler was not called")
273
+ }
274
+
275
+ select {
276
+ case <-done:
277
+ case <-time.After(500 * time.Millisecond):
278
+ t.Fatal("handleConnection did not return after handled no-route callback")
279
+ }
280
+
281
+ _ = client.SetReadDeadline(time.Now().Add(75 * time.Millisecond))
282
+ var b [1]byte
283
+ _, err := client.Read(b[:])
284
+ if err == nil {
285
+ t.Fatal("expected read timeout while connection remains open")
286
+ }
287
+ var netErr net.Error
288
+ if !errors.As(err, &netErr) || !netErr.Timeout() {
289
+ t.Fatalf("expected timeout error to indicate open connection, got: %v", err)
290
+ }
291
+}
292
+
293
+func TestRouter_HandleConnectionNoRouteHandlerDeclined(t *testing.T) {
294
+ router := NewRouter("")
295
+ noRouteCalls := make(chan string, 1)
296
+ router.SetNoRouteHandler(func(_ net.Conn, sni string) bool {
297
+ select {
298
+ case noRouteCalls <- sni:
299
+ default:
300
+ }
301
+ return false
302
+ })
303
+
304
+ client, server := net.Pipe()
305
+ defer func() {
306
+ _ = client.Close()
307
+ _ = server.Close()
308
+ }()
309
+
310
+ done := make(chan struct{})
311
+ router.wg.Add(1)
312
+ go func() {
313
+ router.handleConnection(server)
314
+ close(done)
315
+ }()
316
+
317
+ if _, err := client.Write(buildClientHello("tenant.example.com", true)); err != nil {
318
+ t.Fatalf("write client hello: %v", err)
319
+ }
320
+
321
+ select {
322
+ case <-noRouteCalls:
323
+ case <-time.After(500 * time.Millisecond):
324
+ t.Fatal("no-route handler was not called")
325
+ }
326
+
327
+ select {
328
+ case <-done:
329
+ case <-time.After(500 * time.Millisecond):
330
+ t.Fatal("handleConnection did not return after declined no-route callback")
331
+ }
332
+
333
+ _ = client.SetReadDeadline(time.Now().Add(250 * time.Millisecond))
334
+ var b [1]byte
335
+ _, err := client.Read(b[:])
336
+ if err == nil {
337
+ t.Fatal("expected closed connection after declined no-route callback")
338
+ }
339
+ var netErr net.Error
340
+ if errors.As(err, &netErr) && netErr.Timeout() {
341
+ t.Fatalf("expected closed-connection error, got timeout: %v", err)
342
+ }
343
+}
344
+
345
+func TestCloseWithDebugLogNilCloser(_ *testing.T) {
346
+ closeWithDebugLog(nil, "noop")
347
+}
types/netutil.go
+15
-9
@@ -1,3 +1,4 @@
1
+//nolint:revive // Package name is intentionally aligned with the existing module-wide convention.
2
package types
3
4
import (
@@ -9,6 +10,14 @@ import (
10
"strings"
11
)
12
13
+const defaultBootstrapURL = "http://localhost:4017"
14
+
15
+func normalizeRootHost(raw string) string {
16
+ normalized := strings.ToLower(strings.TrimSpace(raw))
17
+ normalized = strings.TrimPrefix(strings.TrimSuffix(normalized, "."), "*.")
18
+ return normalized
19
+}
20
+
21
// StripScheme removes http:// or https:// prefix from a string.
22
func StripScheme(s string) string {
23
s = strings.TrimSpace(s)
@@ -88,8 +97,7 @@ func parsePortalAddress(raw, fallbackScheme string) (scheme, rootHost, hostPort
97
return "", "", "", false
98
}
99
91
- rootHost = strings.ToLower(strings.TrimSpace(parsed.Hostname()))
92
- rootHost = strings.TrimPrefix(strings.TrimSuffix(rootHost, "."), "*.")
100
+ rootHost = normalizeRootHost(parsed.Hostname())
101
if rootHost == "" {
102
return "", "", "", false
103
}
@@ -187,7 +195,7 @@ func PortalRootHost(portalURL string) string {
195
func DefaultBootstrapFrom(base string) string {
196
base = strings.TrimSpace(base)
197
if base == "" {
190
- return "http://localhost:4017"
198
+ return defaultBootstrapURL
199
}
200
201
if !strings.Contains(base, "://") {
@@ -196,13 +204,13 @@ func DefaultBootstrapFrom(base string) string {
204
205
u, err := url.Parse(strings.TrimSuffix(base, "/"))
206
if err != nil || u.Host == "" {
199
- return "http://localhost:4017"
207
+ return defaultBootstrapURL
208
}
209
if u.Scheme != "http" && u.Scheme != "https" {
202
- return "http://localhost:4017"
210
+ return defaultBootstrapURL
211
}
212
if p := strings.TrimSpace(u.Path); p != "" && p != "/" {
205
- return "http://localhost:4017"
213
+ return defaultBootstrapURL
214
}
215
216
u.Path = ""
@@ -249,9 +257,7 @@ func BuildSNIName(leaseName, baseHost string) string {
257
258
normalizedBaseHost := PortalRootHost(baseHost)
259
if normalizedBaseHost == "" {
252
- normalizedBaseHost = strings.ToLower(strings.TrimSpace(
253
- strings.TrimPrefix(strings.TrimSuffix(baseHost, "."), "*."),
254
- ))
260
+ normalizedBaseHost = normalizeRootHost(baseHost)
261
}
262
if normalizedBaseHost == "" {
263
return ""
types/netutil_test.go
+28
@@ -250,3 +250,31 @@ func TestParsePortalAddressHostOnlyUsesFallbackScheme(t *testing.T) {
250
t.Fatalf("hostPort=%q, want portal.edge.example.com:9443", hostPort)
251
}
252
}
253
+
254
+func TestParsePortalAddressNormalizesTrailingDotAndCase(t *testing.T) {
255
+ t.Parallel()
256
+
257
+ scheme, rootHost, hostPort, ok := parsePortalAddress("HTTPS://Portal.Edge.Example.COM.:7443", "http")
258
+ if !ok {
259
+ t.Fatalf("expected parsePortalAddress success")
260
+ }
261
+ if scheme != "https" {
262
+ t.Fatalf("scheme=%q, want https", scheme)
263
+ }
264
+ if rootHost != "portal.edge.example.com" {
265
+ t.Fatalf("rootHost=%q, want portal.edge.example.com", rootHost)
266
+ }
267
+ if hostPort != "portal.edge.example.com:7443" {
268
+ t.Fatalf("hostPort=%q, want portal.edge.example.com:7443", hostPort)
269
+ }
270
+}
271
+
272
+func TestBuildSNINameFallbackNormalizesRootHost(t *testing.T) {
273
+ t.Parallel()
274
+
275
+ got := BuildSNIName("Api-Gateway", " *.Portal.Edge.Example.COM. ")
276
+ want := "api-gateway.portal.edge.example.com"
277
+ if got != want {
278
+ t.Fatalf("BuildSNIName fallback=%q, want %q", got, want)
279
+ }
280
+}