405 lines
10 KiB
Go
405 lines
10 KiB
Go
package feeds
|
||||
|
|
|
|||
|
|
import (
|
|||
|
|
"context"
|
|||
|
|
"errors"
|
|||
|
|
"fmt"
|
|||
|
|
"io"
|
|||
|
|
"net"
|
|||
|
|
"net/http"
|
|||
|
|
"net/url"
|
|||
|
|
"os"
|
|||
|
|
"strings"
|
|||
|
|
"sync"
|
|||
|
|
"time"
|
|||
|
|
)
|
|||
|
|
|
|||
|
|
const (
|
|||
|
|
// defaultMaxDownloadBytes caps HTTP downloads and local upload sources.
|
|||
|
|
// Bodies stream to a temp file (not RAM), so hundreds of MiB are safe for
|
|||
|
|
// large merchant catalogs. Keep within ~200–500 MiB; raise carefully if disk
|
|||
|
|
// and downloadTimeout remain adequate. Exceeding returns downloadTooLarge().
|
|||
|
|
defaultMaxDownloadBytes int64 = 256 << 20 // 256 MiB
|
|||
|
|
downloadTimeout = 60 * time.Second
|
|||
|
|
dialTimeout = 10 * time.Second
|
|||
|
|
maxRedirects = 5
|
|||
|
|
defaultUserAgent = "DescrybeFeedSync/2.0"
|
|||
|
|
)
|
|||
|
|
|
|||
|
|
// maxDownloadBytes bounds feed downloads (temp file / local source size). Mutable for tests.
|
|||
|
|
var maxDownloadBytes = defaultMaxDownloadBytes
|
|||
|
|
|
|||
|
|
var (
|
|||
|
|
errURLRequired = errors.New("feed url required")
|
|||
|
|
errURLScheme = errors.New("url scheme must be http or https")
|
|||
|
|
errURLPrivate = errors.New("url resolves to a private or blocked address")
|
|||
|
|
errURLFTP = errors.New("ftp/ftps feed sync is not supported yet")
|
|||
|
|
errDownloadTooLarge = errors.New("feed download exceeds size limit")
|
|||
|
|
)
|
|||
|
|
|
|||
|
|
// downloadTooLarge returns errDownloadTooLarge with the active size limit for clients.
|
|||
|
|
func downloadTooLarge() error {
|
|||
|
|
if maxDownloadBytes >= 1<<20 {
|
|||
|
|
return fmt.Errorf("%w (max %d MiB)", errDownloadTooLarge, maxDownloadBytes>>20)
|
|||
|
|
}
|
|||
|
|
return fmt.Errorf("%w (max %d bytes)", errDownloadTooLarge, maxDownloadBytes)
|
|||
|
|
}
|
|||
|
|
|
|||
|
|
var (
|
|||
|
|
allowMu sync.RWMutex
|
|||
|
|
allowConfigured bool
|
|||
|
|
allowHosts map[string]struct{}
|
|||
|
|
allowCIDRs []*net.IPNet
|
|||
|
|
)
|
|||
|
|
|
|||
|
|
// ConfigurePrivateAllowlist sets hostnames and CIDRs that may bypass the private-IP SSRF block.
|
|||
|
|
// Intended for tests and optional startup wiring; production normally uses FEED_URL_PRIVATE_ALLOWLIST
|
|||
|
|
// and/or admin platform setting feeds.private_url_allowlist.
|
|||
|
|
func ConfigurePrivateAllowlist(hosts []string, cidrs []string) error {
|
|||
|
|
h := make(map[string]struct{}, len(hosts))
|
|||
|
|
for _, raw := range hosts {
|
|||
|
|
raw = strings.ToLower(strings.TrimSpace(raw))
|
|||
|
|
if raw != "" {
|
|||
|
|
h[raw] = struct{}{}
|
|||
|
|
}
|
|||
|
|
}
|
|||
|
|
nets := make([]*net.IPNet, 0, len(cidrs))
|
|||
|
|
for _, raw := range cidrs {
|
|||
|
|
raw = strings.TrimSpace(raw)
|
|||
|
|
if raw == "" {
|
|||
|
|
continue
|
|||
|
|
}
|
|||
|
|
_, n, err := net.ParseCIDR(raw)
|
|||
|
|
if err != nil {
|
|||
|
|
return fmt.Errorf("invalid allowlist cidr %q: %w", raw, err)
|
|||
|
|
}
|
|||
|
|
nets = append(nets, n)
|
|||
|
|
}
|
|||
|
|
allowMu.Lock()
|
|||
|
|
allowHosts = h
|
|||
|
|
allowCIDRs = nets
|
|||
|
|
allowConfigured = true
|
|||
|
|
allowMu.Unlock()
|
|||
|
|
return nil
|
|||
|
|
}
|
|||
|
|
|
|||
|
|
// ApplyPrivateAllowlistCSV merges env FEED_URL_PRIVATE_ALLOWLIST with an optional
|
|||
|
|
// admin/settings CSV (settings override by appending unique entries).
|
|||
|
|
func ApplyPrivateAllowlistCSV(settingsCSV string) {
|
|||
|
|
parts := make([]string, 0, 8)
|
|||
|
|
for _, src := range []string{os.Getenv("FEED_URL_PRIVATE_ALLOWLIST"), settingsCSV} {
|
|||
|
|
for _, part := range strings.Split(src, ",") {
|
|||
|
|
part = strings.TrimSpace(part)
|
|||
|
|
if part != "" {
|
|||
|
|
parts = append(parts, part)
|
|||
|
|
}
|
|||
|
|
}
|
|||
|
|
}
|
|||
|
|
hosts := make([]string, 0)
|
|||
|
|
cidrs := make([]string, 0)
|
|||
|
|
seen := map[string]struct{}{}
|
|||
|
|
for _, part := range parts {
|
|||
|
|
key := strings.ToLower(part)
|
|||
|
|
if _, ok := seen[key]; ok {
|
|||
|
|
continue
|
|||
|
|
}
|
|||
|
|
seen[key] = struct{}{}
|
|||
|
|
if strings.Contains(part, "/") {
|
|||
|
|
cidrs = append(cidrs, part)
|
|||
|
|
} else {
|
|||
|
|
hosts = append(hosts, part)
|
|||
|
|
}
|
|||
|
|
}
|
|||
|
|
_ = ConfigurePrivateAllowlist(hosts, cidrs)
|
|||
|
|
}
|
|||
|
|
|
|||
|
|
func ensureAllowlist() {
|
|||
|
|
allowMu.RLock()
|
|||
|
|
done := allowConfigured
|
|||
|
|
allowMu.RUnlock()
|
|||
|
|
if done {
|
|||
|
|
return
|
|||
|
|
}
|
|||
|
|
allowMu.Lock()
|
|||
|
|
defer allowMu.Unlock()
|
|||
|
|
if allowConfigured {
|
|||
|
|
return
|
|||
|
|
}
|
|||
|
|
allowHosts = map[string]struct{}{}
|
|||
|
|
allowCIDRs = nil
|
|||
|
|
raw := strings.TrimSpace(os.Getenv("FEED_URL_PRIVATE_ALLOWLIST"))
|
|||
|
|
for _, part := range strings.Split(raw, ",") {
|
|||
|
|
part = strings.TrimSpace(part)
|
|||
|
|
if part == "" {
|
|||
|
|
continue
|
|||
|
|
}
|
|||
|
|
if strings.Contains(part, "/") {
|
|||
|
|
if _, n, err := net.ParseCIDR(part); err == nil {
|
|||
|
|
allowCIDRs = append(allowCIDRs, n)
|
|||
|
|
}
|
|||
|
|
continue
|
|||
|
|
}
|
|||
|
|
allowHosts[strings.ToLower(part)] = struct{}{}
|
|||
|
|
}
|
|||
|
|
allowConfigured = true
|
|||
|
|
}
|
|||
|
|
|
|||
|
|
// ValidateFeedURL checks a feed source URL for SSRF before persist.
|
|||
|
|
// Empty URL is allowed (local upload sources). Does not fetch the URL.
|
|||
|
|
func ValidateFeedURL(ctx context.Context, rawURL string) error {
|
|||
|
|
ensureAllowlist()
|
|||
|
|
rawURL = strings.TrimSpace(rawURL)
|
|||
|
|
if rawURL == "" {
|
|||
|
|
return nil
|
|||
|
|
}
|
|||
|
|
if len(rawURL) > 2048 {
|
|||
|
|
return errURLScheme
|
|||
|
|
}
|
|||
|
|
lower := strings.ToLower(rawURL)
|
|||
|
|
if strings.HasPrefix(lower, "ftp://") || strings.HasPrefix(lower, "ftps://") {
|
|||
|
|
return errURLFTP
|
|||
|
|
}
|
|||
|
|
u, err := url.Parse(rawURL)
|
|||
|
|
if err != nil || u.Host == "" {
|
|||
|
|
return ClientMsg("invalid url")
|
|||
|
|
}
|
|||
|
|
if u.Scheme != "http" && u.Scheme != "https" {
|
|||
|
|
return errURLScheme
|
|||
|
|
}
|
|||
|
|
return assertPublicHost(ctx, u.Hostname())
|
|||
|
|
}
|
|||
|
|
|
|||
|
|
// downloadFeed streams an http(s) feed to a temp file with SSRF controls and a hard size cap.
|
|||
|
|
// Callers must Close the returned blob to remove the temp file.
|
|||
|
|
func downloadFeed(ctx context.Context, rawURL string) (*feedBlob, error) {
|
|||
|
|
ensureAllowlist()
|
|||
|
|
rawURL = strings.TrimSpace(rawURL)
|
|||
|
|
if rawURL == "" {
|
|||
|
|
return nil, errURLRequired
|
|||
|
|
}
|
|||
|
|
lower := strings.ToLower(rawURL)
|
|||
|
|
if strings.HasPrefix(lower, "ftp://") || strings.HasPrefix(lower, "ftps://") {
|
|||
|
|
return nil, errURLFTP
|
|||
|
|
}
|
|||
|
|
|
|||
|
|
u, err := url.Parse(rawURL)
|
|||
|
|
if err != nil || u.Host == "" {
|
|||
|
|
return nil, ClientMsg("invalid url")
|
|||
|
|
}
|
|||
|
|
if u.Scheme != "http" && u.Scheme != "https" {
|
|||
|
|
return nil, errURLScheme
|
|||
|
|
}
|
|||
|
|
if err := assertPublicHost(ctx, u.Hostname()); err != nil {
|
|||
|
|
return nil, err
|
|||
|
|
}
|
|||
|
|
|
|||
|
|
client := &http.Client{
|
|||
|
|
Timeout: downloadTimeout,
|
|||
|
|
Transport: ssrfTransport(),
|
|||
|
|
CheckRedirect: func(req *http.Request, via []*http.Request) error {
|
|||
|
|
if len(via) >= maxRedirects {
|
|||
|
|
return ClientMsg("too many redirects")
|
|||
|
|
}
|
|||
|
|
if req.URL.Scheme != "http" && req.URL.Scheme != "https" {
|
|||
|
|
return errURLScheme
|
|||
|
|
}
|
|||
|
|
return assertPublicHost(req.Context(), req.URL.Hostname())
|
|||
|
|
},
|
|||
|
|
}
|
|||
|
|
|
|||
|
|
req, err := http.NewRequestWithContext(ctx, http.MethodGet, u.String(), nil)
|
|||
|
|
if err != nil {
|
|||
|
|
return nil, err
|
|||
|
|
}
|
|||
|
|
req.Header.Set("User-Agent", defaultUserAgent)
|
|||
|
|
req.Header.Set("Accept", "text/csv,application/csv,application/xml,text/xml,application/atom+xml,*/*")
|
|||
|
|
|
|||
|
|
res, err := client.Do(req)
|
|||
|
|
if err != nil {
|
|||
|
|
return nil, fmt.Errorf("download failed: %w", err)
|
|||
|
|
}
|
|||
|
|
defer res.Body.Close()
|
|||
|
|
|
|||
|
|
if res.StatusCode < 200 || res.StatusCode >= 300 {
|
|||
|
|
return nil, fmt.Errorf("download status %d", res.StatusCode)
|
|||
|
|
}
|
|||
|
|
if res.ContentLength > maxDownloadBytes {
|
|||
|
|
return nil, downloadTooLarge()
|
|||
|
|
}
|
|||
|
|
|
|||
|
|
tmp, err := os.CreateTemp("", "descrybe-feed-*")
|
|||
|
|
if err != nil {
|
|||
|
|
return nil, fmt.Errorf("temp file: %w", err)
|
|||
|
|
}
|
|||
|
|
tmpPath := tmp.Name()
|
|||
|
|
cleanup := true
|
|||
|
|
defer func() {
|
|||
|
|
_ = tmp.Close()
|
|||
|
|
if cleanup {
|
|||
|
|
_ = os.Remove(tmpPath)
|
|||
|
|
}
|
|||
|
|
}()
|
|||
|
|
|
|||
|
|
limited := io.LimitReader(res.Body, maxDownloadBytes+1)
|
|||
|
|
written, err := io.Copy(tmp, limited)
|
|||
|
|
if err != nil {
|
|||
|
|
return nil, fmt.Errorf("read body: %w", err)
|
|||
|
|
}
|
|||
|
|
if written > maxDownloadBytes {
|
|||
|
|
return nil, downloadTooLarge()
|
|||
|
|
}
|
|||
|
|
if err := tmp.Close(); err != nil {
|
|||
|
|
return nil, fmt.Errorf("close temp: %w", err)
|
|||
|
|
}
|
|||
|
|
cleanup = false
|
|||
|
|
return &feedBlob{
|
|||
|
|
path: tmpPath,
|
|||
|
|
contentType: res.Header.Get("Content-Type"),
|
|||
|
|
size: written,
|
|||
|
|
owned: true,
|
|||
|
|
}, nil
|
|||
|
|
}
|
|||
|
|
|
|||
|
|
func ssrfTransport() *http.Transport {
|
|||
|
|
dialer := &net.Dialer{Timeout: dialTimeout, KeepAlive: 30 * time.Second}
|
|||
|
|
return &http.Transport{
|
|||
|
|
// Never honor HTTP(S)_PROXY: dialing a proxy skips destination SSRF checks.
|
|||
|
|
Proxy: nil,
|
|||
|
|
DialContext: func(ctx context.Context, network, addr string) (net.Conn, error) {
|
|||
|
|
host, port, err := net.SplitHostPort(addr)
|
|||
|
|
if err != nil {
|
|||
|
|
return nil, err
|
|||
|
|
}
|
|||
|
|
if err := assertPublicHost(ctx, host); err != nil {
|
|||
|
|
return nil, err
|
|||
|
|
}
|
|||
|
|
ips, err := resolveHostIPs(ctx, host)
|
|||
|
|
if err != nil {
|
|||
|
|
return nil, err
|
|||
|
|
}
|
|||
|
|
var lastErr error
|
|||
|
|
for _, ip := range ips {
|
|||
|
|
if isBlockedIP(ip) && !isAllowlistedHostOrIP(host, ip) {
|
|||
|
|
lastErr = errURLPrivate
|
|||
|
|
continue
|
|||
|
|
}
|
|||
|
|
target := net.JoinHostPort(ip.String(), port)
|
|||
|
|
conn, err := dialer.DialContext(ctx, network, target)
|
|||
|
|
if err == nil {
|
|||
|
|
return conn, nil
|
|||
|
|
}
|
|||
|
|
lastErr = err
|
|||
|
|
}
|
|||
|
|
if lastErr == nil {
|
|||
|
|
lastErr = errURLPrivate
|
|||
|
|
}
|
|||
|
|
return nil, lastErr
|
|||
|
|
},
|
|||
|
|
ForceAttemptHTTP2: true,
|
|||
|
|
MaxIdleConns: 10,
|
|||
|
|
IdleConnTimeout: 90 * time.Second,
|
|||
|
|
TLSHandshakeTimeout: 10 * time.Second,
|
|||
|
|
ExpectContinueTimeout: 1 * time.Second,
|
|||
|
|
ResponseHeaderTimeout: 30 * time.Second,
|
|||
|
|
}
|
|||
|
|
}
|
|||
|
|
|
|||
|
|
func resolveHostIPs(ctx context.Context, host string) ([]net.IP, error) {
|
|||
|
|
if ip := net.ParseIP(host); ip != nil {
|
|||
|
|
return []net.IP{ip}, nil
|
|||
|
|
}
|
|||
|
|
addrs, err := net.DefaultResolver.LookupIPAddr(ctx, host)
|
|||
|
|
if err != nil {
|
|||
|
|
return nil, fmt.Errorf("dns lookup: %w", err)
|
|||
|
|
}
|
|||
|
|
out := make([]net.IP, 0, len(addrs))
|
|||
|
|
for _, a := range addrs {
|
|||
|
|
out = append(out, a.IP)
|
|||
|
|
}
|
|||
|
|
return out, nil
|
|||
|
|
}
|
|||
|
|
|
|||
|
|
func assertPublicHost(ctx context.Context, host string) error {
|
|||
|
|
ensureAllowlist()
|
|||
|
|
host = strings.TrimSpace(host)
|
|||
|
|
if host == "" {
|
|||
|
|
return errURLPrivate
|
|||
|
|
}
|
|||
|
|
lower := strings.ToLower(host)
|
|||
|
|
allowMu.RLock()
|
|||
|
|
_, hostAllowed := allowHosts[lower]
|
|||
|
|
allowMu.RUnlock()
|
|||
|
|
if hostAllowed {
|
|||
|
|
return nil
|
|||
|
|
}
|
|||
|
|
if lower == "localhost" || strings.HasSuffix(lower, ".localhost") || strings.HasSuffix(lower, ".local") {
|
|||
|
|
return errURLPrivate
|
|||
|
|
}
|
|||
|
|
if ip := net.ParseIP(host); ip != nil {
|
|||
|
|
if isBlockedIP(ip) && !isAllowlistedIP(ip) {
|
|||
|
|
return errURLPrivate
|
|||
|
|
}
|
|||
|
|
return nil
|
|||
|
|
}
|
|||
|
|
|
|||
|
|
addrs, err := resolveHostIPs(ctx, host)
|
|||
|
|
if err != nil {
|
|||
|
|
return err
|
|||
|
|
}
|
|||
|
|
if len(addrs) == 0 {
|
|||
|
|
return errURLPrivate
|
|||
|
|
}
|
|||
|
|
for _, ip := range addrs {
|
|||
|
|
if isBlockedIP(ip) && !isAllowlistedIP(ip) {
|
|||
|
|
return errURLPrivate
|
|||
|
|
}
|
|||
|
|
}
|
|||
|
|
return nil
|
|||
|
|
}
|
|||
|
|
|
|||
|
|
func isAllowlistedHostOrIP(host string, ip net.IP) bool {
|
|||
|
|
ensureAllowlist()
|
|||
|
|
allowMu.RLock()
|
|||
|
|
_, ok := allowHosts[strings.ToLower(strings.TrimSpace(host))]
|
|||
|
|
allowMu.RUnlock()
|
|||
|
|
if ok {
|
|||
|
|
return true
|
|||
|
|
}
|
|||
|
|
return isAllowlistedIP(ip)
|
|||
|
|
}
|
|||
|
|
|
|||
|
|
func isAllowlistedIP(ip net.IP) bool {
|
|||
|
|
ensureAllowlist()
|
|||
|
|
if ip == nil {
|
|||
|
|
return false
|
|||
|
|
}
|
|||
|
|
allowMu.RLock()
|
|||
|
|
defer allowMu.RUnlock()
|
|||
|
|
for _, n := range allowCIDRs {
|
|||
|
|
if n.Contains(ip) {
|
|||
|
|
return true
|
|||
|
|
}
|
|||
|
|
}
|
|||
|
|
return false
|
|||
|
|
}
|
|||
|
|
|
|||
|
|
func isBlockedIP(ip net.IP) bool {
|
|||
|
|
if ip == nil {
|
|||
|
|
return true
|
|||
|
|
}
|
|||
|
|
if ip.IsLoopback() || ip.IsPrivate() || ip.IsLinkLocalUnicast() || ip.IsLinkLocalMulticast() ||
|
|||
|
|
ip.IsMulticast() || ip.IsUnspecified() {
|
|||
|
|
return true
|
|||
|
|
}
|
|||
|
|
// AWS/GCP/Azure metadata and CGNAT.
|
|||
|
|
if ip4 := ip.To4(); ip4 != nil {
|
|||
|
|
if ip4[0] == 169 && ip4[1] == 254 {
|
|||
|
|
return true
|
|||
|
|
}
|
|||
|
|
if ip4[0] == 100 && ip4[1] >= 64 && ip4[1] <= 127 {
|
|||
|
|
return true
|
|||
|
|
}
|
|||
|
|
}
|
|||
|
|
return false
|
|||
|
|
}
|