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(&registerReq); 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, &registerReq, "[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 +}