Files
descrybe/apps/api/internal/feeds/download.go
T

405 lines
10 KiB
Go
Raw Normal View History

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 ~200500 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
}