From 276d767ce5baad29fe66d115a23a2d8b7fc4cab8 Mon Sep 17 00:00:00 2001 From: =?UTF-8?q?Deluan=20Quint=C3=A3o?= Date: Sat, 12 Sep 2026 13:41:00 -0400 Subject: [PATCH] Merge commit from fork * fix(plugins): apply the private-address dial guard to WebSocket connections The WebSocket host service only matched the host string against requiredHosts, so an allowlisted name resolving (or rebinding) to a private address was dialed. Share the HTTP client's resolved-IP check and allowlist matching, so WebSocket follows the same rules: named hosts can't reach private addresses, literal IP/CIDR entries and a bare "*" can. * refactor(plugins): drop redundant WebSocket dial timeout and tidy guard tests --- plugins/host_httpclient.go | 48 ++-------------------------- plugins/host_netguard.go | 58 ++++++++++++++++++++++++++++++++++ plugins/host_websocket.go | 14 ++++---- plugins/host_websocket_test.go | 45 ++++++++++++++++++++++++++ 4 files changed, 112 insertions(+), 53 deletions(-) create mode 100644 plugins/host_netguard.go diff --git a/plugins/host_httpclient.go b/plugins/host_httpclient.go index 749ba096e..209f80eb0 100644 --- a/plugins/host_httpclient.go +++ b/plugins/host_httpclient.go @@ -9,7 +9,6 @@ import ( "net" "net/http" "net/url" - "slices" "strings" "syscall" "time" @@ -174,51 +173,12 @@ func (s *httpServiceImpl) validateHost(ctx context.Context, hostStr string) erro return nil } -// dialControl checks the resolved IP, so hostnames can't reach private addresses unless a literal -// IP/CIDR entry or a bare "*" (plugins targeting user-configured LAN services) allows it. func (s *httpServiceImpl) dialControl(_, address string, _ syscall.RawConn) error { - if slices.Contains(s.requiredHosts, "*") { - return nil - } - host, _, err := net.SplitHostPort(address) - if err != nil { - return err - } - ip := net.ParseIP(host) - if ip == nil || !isPrivateIP(ip) { - return nil - } - for _, entry := range s.requiredHosts { - if ipMatchesEntry(entry, ip) { - return nil - } - } - return fmt.Errorf("dial to private/loopback address %q blocked: requires an explicit IP or CIDR in requiredHosts", address) -} - -// ipMatchesEntry reports whether a requiredHosts entry is a literal IP or CIDR -// that covers ip. Hostname and wildcard entries never match. -func ipMatchesEntry(entry string, ip net.IP) bool { - if _, cidr, err := net.ParseCIDR(entry); err == nil { - return cidr.Contains(ip) - } - if entryIP := net.ParseIP(entry); entryIP != nil { - return entryIP.Equal(ip) - } - return false + return checkPrivateDial(s.requiredHosts, address) } func (s *httpServiceImpl) isHostAllowed(hostname string) bool { - ip := net.ParseIP(hostname) - for _, pattern := range s.requiredHosts { - if matchHostPattern(pattern, hostname) { - return true - } - if ip != nil && ipMatchesEntry(pattern, ip) { - return true - } - } - return false + return isHostInAllowlist(s.requiredHosts, hostname) } // extractHostname returns the hostname portion of a host string, stripping @@ -248,9 +208,5 @@ func isPrivateOrLoopback(hostname string) bool { return isPrivateIP(ip) } -func isPrivateIP(ip net.IP) bool { - return ip.IsLoopback() || ip.IsUnspecified() || ip.IsPrivate() || ip.IsLinkLocalUnicast() || ip.IsLinkLocalMulticast() -} - // Verify interface implementation var _ host.HTTPService = (*httpServiceImpl)(nil) diff --git a/plugins/host_netguard.go b/plugins/host_netguard.go new file mode 100644 index 000000000..8400446ed --- /dev/null +++ b/plugins/host_netguard.go @@ -0,0 +1,58 @@ +package plugins + +import ( + "fmt" + "net" + "slices" +) + +// checkPrivateDial runs at dial time on the resolved IP, so hostnames can't reach private addresses unless a +// literal IP/CIDR entry or a bare "*" (plugins targeting user-configured LAN services) allows it. +func checkPrivateDial(requiredHosts []string, address string) error { + if slices.Contains(requiredHosts, "*") { + return nil + } + host, _, err := net.SplitHostPort(address) + if err != nil { + return err + } + ip := net.ParseIP(host) + if ip == nil || !isPrivateIP(ip) { + return nil + } + for _, entry := range requiredHosts { + if ipMatchesEntry(entry, ip) { + return nil + } + } + return fmt.Errorf("dial to private/loopback address %q blocked: requires an explicit IP or CIDR in requiredHosts", address) +} + +func isHostInAllowlist(requiredHosts []string, hostname string) bool { + ip := net.ParseIP(hostname) + for _, pattern := range requiredHosts { + if matchHostPattern(pattern, hostname) { + return true + } + if ip != nil && ipMatchesEntry(pattern, ip) { + return true + } + } + return false +} + +// ipMatchesEntry reports whether a requiredHosts entry is a literal IP or CIDR +// that covers ip. Hostname and wildcard entries never match. +func ipMatchesEntry(entry string, ip net.IP) bool { + if _, cidr, err := net.ParseCIDR(entry); err == nil { + return cidr.Contains(ip) + } + if entryIP := net.ParseIP(entry); entryIP != nil { + return entryIP.Equal(ip) + } + return false +} + +func isPrivateIP(ip net.IP) bool { + return ip.IsLoopback() || ip.IsUnspecified() || ip.IsPrivate() || ip.IsLinkLocalUnicast() || ip.IsLinkLocalMulticast() +} diff --git a/plugins/host_websocket.go b/plugins/host_websocket.go index 82aded0cb..933e53144 100644 --- a/plugins/host_websocket.go +++ b/plugins/host_websocket.go @@ -5,10 +5,12 @@ import ( "errors" "fmt" "maps" + "net" "net/http" "net/url" "strings" "sync" + "syscall" "time" "github.com/gorilla/websocket" @@ -112,6 +114,7 @@ func (s *webSocketServiceImpl) Connect(ctx context.Context, urlStr string, heade // Establish WebSocket connection dialer := websocket.Dialer{ HandshakeTimeout: 30 * time.Second, + NetDialContext: (&net.Dialer{Control: s.dialControl}).DialContext, } conn, resp, err := dialer.DialContext(ctx, urlStr, httpHeaders) @@ -243,14 +246,11 @@ func (s *webSocketServiceImpl) getConnection(connectionID string) (*wsConnection } func (s *webSocketServiceImpl) isHostAllowed(host string) bool { - hostWithoutPort := extractHostname(host) + return isHostInAllowlist(s.requiredHosts, extractHostname(host)) +} - for _, pattern := range s.requiredHosts { - if matchHostPattern(pattern, hostWithoutPort) { - return true - } - } - return false +func (s *webSocketServiceImpl) dialControl(_, address string, _ syscall.RawConn) error { + return checkPrivateDial(s.requiredHosts, address) } // matchHostPattern matches a host against a pattern. diff --git a/plugins/host_websocket_test.go b/plugins/host_websocket_test.go index 9f2d20bef..772d83cc9 100644 --- a/plugins/host_websocket_test.go +++ b/plugins/host_websocket_test.go @@ -8,6 +8,7 @@ import ( "maps" "net/http" "net/http/httptest" + "net/url" "os" "path/filepath" "strings" @@ -499,6 +500,50 @@ var _ = Describe("WebSocketService", Ordered, func() { }) }) + Describe("Private address protection", func() { + var wsServer *httptest.Server + var savedHosts []string + + BeforeEach(func() { + savedHosts = testService.requiredHosts + upgrader := websocket.Upgrader{CheckOrigin: func(r *http.Request) bool { return true }} + wsServer = httptest.NewServer(http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) { + if conn, err := upgrader.Upgrade(w, r, nil); err == nil { + _, _, _ = conn.ReadMessage() + } + })) + }) + + AfterEach(func() { + testService.closeAllConnections() + testService.requiredHosts = savedHosts + wsServer.Close() + }) + + serverPort := func() string { + u, _ := url.Parse(wsServer.URL) + return u.Port() + } + + It("blocks an allowlisted hostname that resolves to loopback", func() { + testService.requiredHosts = []string{"localhost."} + _, err := testService.Connect(GinkgoT().Context(), "ws://localhost.:"+serverPort(), nil, "") + Expect(err).To(MatchError(ContainSubstring("private/loopback"))) + }) + + It("allows loopback when a CIDR entry covers it", func() { + testService.requiredHosts = []string{"127.0.0.0/8"} + _, err := testService.Connect(GinkgoT().Context(), "ws://127.0.0.1:"+serverPort(), nil, "") + Expect(err).ToNot(HaveOccurred()) + }) + + It("allows loopback when the allowlist is the bare '*' wildcard", func() { + testService.requiredHosts = []string{"*"} + _, err := testService.Connect(GinkgoT().Context(), "ws://localhost.:"+serverPort(), nil, "") + Expect(err).ToNot(HaveOccurred()) + }) + }) + Describe("Plugin Unload", func() { It("should close all connections when plugin is unloaded", func() { // Create a fresh server for this test