mirror of
https://github.com/navidrome/navidrome.git
synced 2026-10-08 02:17:25 +02:00
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:
parent
1a8463f7de
commit
276d767ce5
4 changed files with 112 additions and 53 deletions
|
|
@ -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
58
plugins/host_netguard.go
Normal 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()
|
||||
}
|
||||
|
|
@ -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.
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
|
|
|||
Loading…
Add table
Add a link
Reference in a new issue