diff --git a/appview/serververify/verify.go b/appview/serververify/verify.go index f7d2a466..77071933 100644 --- a/appview/serververify/verify.go +++ b/appview/serververify/verify.go @@ -4,9 +4,7 @@ import ( "context" "errors" "fmt" - "net" "net/http" - "syscall" "time" indigoxrpc "github.com/bluesky-social/indigo/xrpc" @@ -14,6 +12,7 @@ import ( "tangled.org/core/appview/db" "tangled.org/core/orm" "tangled.org/core/rbac" + "tangled.org/core/util/netutil" "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/knotmirror/knotstream/slurper.go b/knotmirror/knotstream/slurper.go index 36f4318d..08b5c08f 100644 --- a/knotmirror/knotstream/slurper.go +++ b/knotmirror/knotstream/slurper.go @@ -12,13 +12,13 @@ import ( "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/util/netutil" ) type KnotSlurper struct { @@ -135,8 +135,7 @@ func (s *KnotSlurper) subscribeWithRedialer(ctx context.Context, host models.Hos // 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 index f6d308bd..21ab8f5c 100644 --- a/knotmirror/xrpc/xrpc.go +++ b/knotmirror/xrpc/xrpc.go @@ -9,7 +9,6 @@ import ( "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 @@ import ( "tangled.org/core/knotmirror/knotstream" "tangled.org/core/knotmirror/repoindexer" "tangled.org/core/log" + "tangled.org/core/util/netutil" ) type Xrpc struct { @@ -37,7 +37,7 @@ func New(logger *slog.Logger, cfg *config.Config, db *sql.DB, rdb *redis.Client, Timeout: 30 * time.Second, } if cfg.KnotSSRF { - httpClient.Transport = ssrf.PublicOnlyTransport() + httpClient.Transport = netutil.SSRFTransport(false) } return &Xrpc{ cfg: cfg, diff --git a/spindle/mill/executor/executor.go b/spindle/mill/executor/executor.go index 583ecec2..16b187b9 100644 --- a/spindle/mill/executor/executor.go +++ b/spindle/mill/executor/executor.go @@ -14,7 +14,6 @@ import ( "time" "github.com/bluesky-social/indigo/atproto/syntax" - "github.com/gorilla/websocket" "tangled.org/core/api/tangled" "tangled.org/core/notifier" @@ -22,6 +21,7 @@ import ( "tangled.org/core/spindle/db" "tangled.org/core/spindle/engine" "tangled.org/core/spindle/models" + "tangled.org/core/util/netutil" millproto "tangled.org/core/spindle/mill/proto" millv1 "tangled.org/core/spindle/mill/proto/gen" @@ -131,11 +131,15 @@ func (e *Executor) Connect(ctx context.Context) { } func (e *Executor) runSession(ctx context.Context) error { + dev := e.cfg.Server.Dev + if _, err := netutil.EnforceWSSURL(e.millURL, dev); err != nil { + return fmt.Errorf("mill url: %w", err) + } header := http.Header{} if e.token != "" { header.Set("Authorization", "Bearer "+e.token) } - conn, _, err := websocket.DefaultDialer.DialContext(ctx, e.millURL, header) + conn, _, err := netutil.SSRFWebsocketDialer(dev).DialContext(ctx, e.millURL, header) if err != nil { return fmt.Errorf("dial mill: %w", err) } diff --git a/spindle/mill/integration_test.go b/spindle/mill/integration_test.go index 8413d6a4..bbdf4b0d 100644 --- a/spindle/mill/integration_test.go +++ b/spindle/mill/integration_test.go @@ -49,6 +49,7 @@ func TestEndToEndDummyJob(t *testing.T) { } en := notifier.New() cfg := &config.Config{} + cfg.Server.Dev = true cfg.Server.LogDir = execDir cfg.Server.Hostname = "exec-1" cfg.Mill.URL = wsURL @@ -116,6 +117,7 @@ func TestExecutorConfiguredLabelsAreStoredOnSession(t *testing.T) { } en := notifier.New() cfg := &config.Config{} + cfg.Server.Dev = true cfg.Server.LogDir = execDir cfg.Server.Hostname = "exec-labels" cfg.Mill.URL = wsURL @@ -165,6 +167,7 @@ func TestEndToEndDummyJobUsesRequiredLabelsAcrossExecutors(t *testing.T) { } en := notifier.New() cfg := &config.Config{} + cfg.Server.Dev = true cfg.Server.LogDir = execDir cfg.Server.Hostname = name cfg.Mill.URL = wsURL diff --git a/util/netutil/ssrf.go b/util/netutil/ssrf.go new file mode 100644 index 00000000..021a7e79 --- /dev/null +++ b/util/netutil/ssrf.go @@ -0,0 +1,52 @@ +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 { + return &websocket.Dialer{ + Proxy: http.ProxyFromEnvironment, + NetDialContext: SSRFDialer(dev).DialContext, + } +} + +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 +}