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
This commit is contained in:
Deluan Quintão 2026-09-12 13:41:00 -04:00 • committed by GitHub
commit 276d767ce5
No known key found for this signature in database
GPG key ID: B5690EEEBB952194
4 changed files with 112 additions and 53 deletions

View file

@ -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)

58
plugins/host_netguard.go Normal file
View file

@ -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()
}

View file

@ -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.

View file

@ -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