diff --git a/netutil/ssrf.go b/netutil/ssrf.go new file mode 100644 --- /dev/null +++ b/netutil/ssrf.go @@ -0,0 +1,51 @@ +package netutil + +import ( + "fmt" + "net" + "net/http" + "net/url" + + "github.com/bluesky-social/indigo/util/ssrf" + "github.com/gorilla/websocket" +) + +// SSRFDialer returns a net.Dialer that refuses non-public IPs. +func SSRFDialer(dev bool) *net.Dialer { + if dev { + return &net.Dialer{} + } + return ssrf.PublicOnlyDialer() +} + +// SSRFTransport returns an http.Transport that refuses non-public IPs. +func SSRFTransport(dev bool) *http.Transport { + if dev { + return &http.Transport{} + } + return ssrf.PublicOnlyTransport() +} + +// SSRFWebsocketDialer returns a websocket.Dialer that refuses non-public IPs. +func SSRFWebsocketDialer(dev bool) *websocket.Dialer { + dialer := *websocket.DefaultDialer + dialer.NetDialContext = SSRFDialer(dev).DialContext + return &dialer +} + +func EnforceWSSURL(rawURL string, dev bool) (*url.URL, error) { + u, err := url.Parse(rawURL) + if err != nil { + return nil, fmt.Errorf("invalid url: %w", err) + } + switch u.Scheme { + case "wss": + case "ws": + if !dev { + return nil, fmt.Errorf("insecure scheme %q is prohibited in production; use wss://", u.Scheme) + } + default: + return nil, fmt.Errorf("unsupported websocket scheme %q", u.Scheme) + } + return u, nil +} diff --git a/netutil/ssrf_test.go b/netutil/ssrf_test.go new file mode 100644 --- /dev/null +++ b/netutil/ssrf_test.go @@ -0,0 +1,17 @@ +package netutil + +import ( + "testing" + + "github.com/gorilla/websocket" +) + +func TestSSRFWebsocketDialerPreservesHandshakeTimeout(t *testing.T) { + dialer := SSRFWebsocketDialer(false) + if dialer.HandshakeTimeout != websocket.DefaultDialer.HandshakeTimeout { + t.Fatalf("HandshakeTimeout = %v, want %v", dialer.HandshakeTimeout, websocket.DefaultDialer.HandshakeTimeout) + } + if dialer.NetDialContext == nil { + t.Fatal("NetDialContext is nil; public-only dialing is not enforced") + } +} diff --git a/appview/serververify/verify.go b/appview/serververify/verify.go --- a/appview/serververify/verify.go +++ b/appview/serververify/verify.go @@ -4,14 +4,13 @@ "context" "errors" "fmt" - "net" "net/http" - "syscall" "time" indigoxrpc "github.com/bluesky-social/indigo/xrpc" "tangled.org/core/api/tangled" "tangled.org/core/appview/db" + "tangled.org/core/netutil" "tangled.org/core/orm" "tangled.org/core/rbac" "tangled.org/core/xrpc/xrpcclient" @@ -31,8 +30,12 @@ } host := fmt.Sprintf("%s://%s", scheme, domain) + dialer := netutil.SSRFDialer(dev) + dialer.Timeout = 5 * time.Second + dialer.KeepAlive = 30 * time.Second + transport := &http.Transport{ - DialContext: safeDialer(dev).DialContext, + DialContext: dialer.DialContext, } xrpcc := &indigoxrpc.Client{ Host: host, @@ -175,29 +178,4 @@ committed = true return nil -} -func safeDialer(dev bool) *net.Dialer { - d := &net.Dialer{ - Timeout: 5 * time.Second, - KeepAlive: 30 * time.Second, - } - if dev { - return d - } - d.Control = func(network, address string, _ syscall.RawConn) error { - host, _, err := net.SplitHostPort(address) - if err != nil { - return fmt.Errorf("invalid dial address %q: %w", address, err) - } - ip := net.ParseIP(host) - if ip == nil { - return fmt.Errorf("dial address %q did not resolve to IP", address) - } - if ip.IsLoopback() || ip.IsPrivate() || ip.IsLinkLocalUnicast() || - ip.IsLinkLocalMulticast() || ip.IsMulticast() || ip.IsUnspecified() { - return fmt.Errorf("refusing to dial %s: reserved or private address", ip) - } - return nil - } - return d } diff --git a/knotmirror/knotstream/slurper.go b/knotmirror/knotstream/slurper.go --- a/knotmirror/knotstream/slurper.go +++ b/knotmirror/knotstream/slurper.go @@ -12,13 +12,13 @@ "time" "github.com/bluesky-social/indigo/atproto/syntax" - "github.com/bluesky-social/indigo/util/ssrf" "github.com/carlmjohnson/versioninfo" "github.com/gorilla/websocket" "tangled.org/core/knotmirror/config" "tangled.org/core/knotmirror/db" "tangled.org/core/knotmirror/models" "tangled.org/core/log" + "tangled.org/core/netutil" ) type KnotSlurper struct { @@ -135,8 +135,7 @@ // if this isn't a localhost / private connection, then we should enable SSRF protections if !host.NoSSL || s.ssrf { - netDialer := ssrf.PublicOnlyDialer() - dialer.NetDialContext = netDialer.DialContext + dialer.NetDialContext = netutil.SSRFDialer(false).DialContext } cursor := host.LastSeq diff --git a/knotmirror/xrpc/xrpc.go b/knotmirror/xrpc/xrpc.go --- a/knotmirror/xrpc/xrpc.go +++ b/knotmirror/xrpc/xrpc.go @@ -9,7 +9,6 @@ "time" "github.com/bluesky-social/indigo/atproto/atclient" - "github.com/bluesky-social/indigo/util/ssrf" "github.com/go-chi/chi/v5" "github.com/redis/go-redis/v9" "tangled.org/core/api/tangled" @@ -18,6 +17,7 @@ "tangled.org/core/knotmirror/knotstream" "tangled.org/core/knotmirror/repoindexer" "tangled.org/core/log" + "tangled.org/core/netutil" ) type Xrpc struct { @@ -37,7 +37,7 @@ Timeout: 30 * time.Second, } if cfg.KnotSSRF { - httpClient.Transport = ssrf.PublicOnlyTransport() + httpClient.Transport = netutil.SSRFTransport(false) } return &Xrpc{ cfg: cfg,