From cf2654b85d04ffda1cd288236553b8b407153e82 Mon Sep 17 00:00:00 2001 From: dawn Date: Sun, 19 Jul 2026 02:50:28 +0300 Subject: [PATCH] netutil: extract common SSRF and secure-scheme guards Signed-off-by: dawn --- appview/serververify/verify.go | 34 ++++------------------- netutil/ssrf.go | 51 ++++++++++++++++++++++++++++++++++ netutil/ssrf_test.go | 17 ++++++++++++ 3 files changed, 74 insertions(+), 28 deletions(-) create mode 100644 netutil/ssrf.go create mode 100644 netutil/ssrf_test.go diff --git a/appview/serververify/verify.go b/appview/serververify/verify.go index f7d2a466..7c97c73d 100644 --- a/appview/serververify/verify.go +++ b/appview/serververify/verify.go @@ -4,14 +4,13 @@ import ( "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 @@ func fetchOwner(ctx context.Context, domain string, dev bool) (string, error) { } 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, @@ -176,28 +179,3 @@ func MarkKnotVerified(d *db.DB, e *rbac.Enforcer, domain, owner string) error { 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/netutil/ssrf.go b/netutil/ssrf.go new file mode 100644 index 00000000..6974187d --- /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" +) + +// refuses non-public ips to prevent ssrf +func SSRFDialer(dev bool) *net.Dialer { + if dev { + return &net.Dialer{} + } + return ssrf.PublicOnlyDialer() +} + +// refuses non-public ips to prevent ssrf +func SSRFTransport(dev bool) *http.Transport { + if dev { + return &http.Transport{} + } + return ssrf.PublicOnlyTransport() +} + +// refuses non-public ips to prevent ssrf +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 index 00000000..8ca51fe7 --- /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") + } +} -- 2.51.2