diff --git a/knotmirror/knotstream/knotstream.go b/knotmirror/knotstream/knotstream.go --- a/knotmirror/knotstream/knotstream.go +++ b/knotmirror/knotstream/knotstream.go @@ -24,7 +24,7 @@ l = log.SubLogger(l, "knotstream") return &KnotStream{ logger: l, db: db, - slurper: NewKnotSlurper(l, db, cfg.Slurper), + slurper: NewKnotSlurper(l, db, cfg), } } diff --git a/knotmirror/knotstream/slurper.go b/knotmirror/knotstream/slurper.go --- a/knotmirror/knotstream/slurper.go +++ b/knotmirror/knotstream/slurper.go @@ -25,16 +25,18 @@ type KnotSlurper struct { logger *slog.Logger db *sql.DB cfg config.SlurperConfig + ssrf bool subsLk sync.Mutex subs map[string]*subscription } -func NewKnotSlurper(l *slog.Logger, db *sql.DB, cfg config.SlurperConfig) *KnotSlurper { +func NewKnotSlurper(l *slog.Logger, db *sql.DB, cfg *config.Config) *KnotSlurper { return &KnotSlurper{ logger: log.SubLogger(l, "slurper"), db: db, - cfg: cfg, + cfg: cfg.Slurper, + ssrf: cfg.KnotSSRF, subs: make(map[string]*subscription), } } @@ -132,7 +134,7 @@ HandshakeTimeout: time.Second * 5, } // if this isn't a localhost / private connection, then we should enable SSRF protections - if !host.NoSSL { + if !host.NoSSL || s.ssrf { netDialer := ssrf.PublicOnlyDialer() dialer.NetDialContext = netDialer.DialContext } diff --git a/knotmirror/xrpc/proxy.go b/knotmirror/xrpc/proxy.go --- a/knotmirror/xrpc/proxy.go +++ b/knotmirror/xrpc/proxy.go @@ -2,6 +2,7 @@ package xrpc import ( "context" + "errors" "fmt" "io" "net/http" @@ -45,6 +46,30 @@ baseURL string repoIdentifier string } +// validateKnotURL ensures a knot base URL is safe to proxy to. +// It rejects URLs with path components, query strings, or fragments +// that could be used for path injection. +func validateKnotURL(raw string) (string, error) { + u, err := url.Parse(raw) + if err != nil { + return "", fmt.Errorf("invalid knot URL: %w", err) + } + if u.Scheme != "http" && u.Scheme != "https" { + return "", errors.New("knot URL must use http or https scheme") + } + if u.Path != "" && u.Path != "/" { + return "", fmt.Errorf("knot URL must not contain a path: %q", raw) + } + if u.RawQuery != "" || u.Fragment != "" { + return "", fmt.Errorf("knot URL must not contain query or fragment: %q", raw) + } + if u.User != nil { + return "", fmt.Errorf("knot URL must not contain userinfo: %q", raw) + } + // Strip trailing slash for consistent formatting + return strings.TrimRight(u.String(), "/"), nil +} + func (x *Xrpc) resolveKnot(ctx context.Context, repoAt syntax.ATURI) (*knotInfo, error) { repo, err := db.GetRepoByAtUri(ctx, x.db, repoAt) if err == nil && repo != nil { @@ -60,6 +85,10 @@ } else { knotURL = "http://" + knotURL } } + } + knotURL, err = validateKnotURL(knotURL) + if err != nil { + return nil, err } return &knotInfo{baseURL: knotURL, repoIdentifier: repo.RepoIdentifier()}, nil } @@ -111,6 +140,10 @@ x.logger.Error("failed to upsert repo after proxy resolution", "err", upsertErr) } }() + knotURL, err = validateKnotURL(knotURL) + if err != nil { + return nil, err + } return &knotInfo{ baseURL: knotURL, repoIdentifier: repoDid.String(), diff --git a/knotmirror/xrpc/xrpc.go b/knotmirror/xrpc/xrpc.go --- a/knotmirror/xrpc/xrpc.go +++ b/knotmirror/xrpc/xrpc.go @@ -9,6 +9,7 @@ "net/http" "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" @@ -30,17 +31,21 @@ inflight *inflightTracker } func New(logger *slog.Logger, cfg *config.Config, db *sql.DB, rdb *redis.Client, resolver *idresolver.Resolver, ks *knotstream.KnotStream) *Xrpc { + httpClient := &http.Client{ + Timeout: 30 * time.Second, + } + if cfg.KnotSSRF { + httpClient.Transport = ssrf.PublicOnlyTransport() + } return &Xrpc{ - cfg: cfg, - db: db, - rdb: rdb, - resolver: resolver, - ks: ks, - logger: log.SubLogger(logger, "xrpc"), - httpClient: &http.Client{ - Timeout: 30 * time.Second, - }, - inflight: newInflightTracker(), + cfg: cfg, + db: db, + rdb: rdb, + resolver: resolver, + ks: ks, + logger: log.SubLogger(logger, "xrpc"), + httpClient: httpClient, + inflight: newInflightTracker(), } }