diff --git a/fetchguard/fetch.go b/fetchguard/fetch.go new file mode 100644 index 0000000..fd1d81d --- /dev/null +++ b/fetchguard/fetch.go @@ -0,0 +1,175 @@ +package fetchguard + +import ( + "context" + "errors" + "fmt" + "io" + "net" + "net/http" + "net/netip" + "net/url" + "strings" + "syscall" + "time" + + "git.golem15.com/golem15/summercms/compass" +) + +// Result is a successful (including non-2xx, including 3xx) Fetch response. +type Result struct { + Body []byte + ContentType string + StatusCode int +} + +var errPrivateIP = errors.New("private_ip") + +// Fetch validates url against policy, resolves defaults for any zero +// MaxBytes/Timeout via DefaultsFromConfig (or Defaults() when cfg is nil), +// then performs the guarded HTTPS GET. +// +// Redirects are never followed: CheckRedirect returns http.ErrUseLastResponse, +// so a 3xx response is returned as a non-error Result. Callers that want to +// follow Location must re-invoke Fetch, which re-runs the same guard. +// +// A non-nil error is always *Error with a Reason from the closed set. +func Fetch(ctx context.Context, rawURL string, policy Policy, cfg *compass.Config) (*Result, error) { + if ctx == nil { + ctx = context.Background() + } + parsed, err := url.Parse(rawURL) + if err != nil || parsed.Scheme == "" || parsed.Host == "" { + return nil, &Error{Reason: ReasonInvalidURL, Err: err} + } + if strings.ToLower(parsed.Scheme) != "https" { + return nil, &Error{Reason: ReasonScheme} + } + host := parsed.Hostname() + if policy.Mode == AllowHostsMode && !hostAllowed(host, policy.AllowHosts) { + return nil, &Error{Reason: ReasonInvalidURL} + } + + maxBytes, timeout, err := resolveLimits(policy, cfg) + if err != nil { + return nil, err + } + + client := &http.Client{ + Timeout: timeout, + CheckRedirect: func(*http.Request, []*http.Request) error { + return http.ErrUseLastResponse + }, + Transport: &http.Transport{ + // User-supplied URLs must not be forwarded through HTTP_PROXY: + // the dial-time IP check would then see the proxy, not the target. + Proxy: nil, + DialContext: (&net.Dialer{ + Timeout: timeout, + Control: dialControl(policy), + }).DialContext, + TLSClientConfig: policy.tlsConfig, + DisableKeepAlives: true, + ForceAttemptHTTP2: true, + }, + } + + req, err := http.NewRequestWithContext(ctx, http.MethodGet, parsed.String(), nil) + if err != nil { + return nil, &Error{Reason: ReasonInvalidURL, Err: err} + } + resp, err := client.Do(req) + if err != nil { + return nil, mapTransportError(err) + } + defer resp.Body.Close() + + data, err := io.ReadAll(io.LimitReader(resp.Body, maxBytes+1)) + if err != nil { + return nil, &Error{Reason: ReasonNetworkError, Err: err} + } + if int64(len(data)) == maxBytes+1 { + return nil, &Error{Reason: ReasonTooLarge} + } + return &Result{ + Body: data, + ContentType: resp.Header.Get("Content-Type"), + StatusCode: resp.StatusCode, + }, nil +} + +func resolveLimits(policy Policy, cfg *compass.Config) (int64, time.Duration, error) { + maxBytes := policy.MaxBytes + timeout := policy.Timeout + if maxBytes <= 0 || timeout <= 0 { + var ( + defMax int64 + defTO time.Duration + err error + ) + if cfg != nil { + defMax, defTO, err = DefaultsFromConfig(cfg) + if err != nil { + return 0, 0, &Error{Reason: ReasonInvalidURL, Err: err} + } + } else { + defMax, defTO = Defaults() + } + if maxBytes <= 0 { + maxBytes = defMax + } + if timeout <= 0 { + timeout = defTO + } + } + if maxBytes <= 0 || timeout <= 0 { + return 0, 0, &Error{Reason: ReasonInvalidURL, Err: fmt.Errorf("max bytes and timeout must be positive")} + } + return maxBytes, timeout, nil +} + +func hostAllowed(host string, allowed []string) bool { + host = strings.ToLower(host) + for _, a := range allowed { + a = strings.ToLower(a) + if a == "" { + continue + } + if host == a || strings.HasSuffix(host, "."+a) { + return true + } + } + return false +} + +func dialControl(policy Policy) func(network, address string, c syscall.RawConn) error { + return func(network, address string, c syscall.RawConn) error { + if policy.skipReservedCheck { + return nil + } + host, _, err := net.SplitHostPort(address) + if err != nil { + return err + } + addr, err := netip.ParseAddr(host) + if err != nil { + return fmt.Errorf("fetchguard: unparseable dial address %q: %w", host, err) + } + addr = addr.Unmap() + if isReservedOrPrivate(addr) { + return errPrivateIP + } + return nil + } +} + +func mapTransportError(err error) *Error { + if errors.Is(err, errPrivateIP) { + return &Error{Reason: ReasonPrivateIP, Err: err} + } + var dnsErr *net.DNSError + if errors.As(err, &dnsErr) { + return &Error{Reason: ReasonUnresolvable, Err: err} + } + return &Error{Reason: ReasonNetworkError, Err: err} +} diff --git a/fetchguard/policy.go b/fetchguard/policy.go new file mode 100644 index 0000000..b861010 --- /dev/null +++ b/fetchguard/policy.go @@ -0,0 +1,147 @@ +package fetchguard + +import ( + "crypto/tls" + "fmt" + "math" + "time" + + "git.golem15.com/golem15/summercms/compass" +) + +// Mode selects host-allow-list vs any-public-host. The private/loopback/ +// reserved IP block is always on regardless of Mode (D-11). +type Mode int + +const ( + AllowHostsMode Mode = iota + PublicOnlyMode +) + +// Reason is the closed set of Fetch failure reasons, matching PHP's +// invalid_url/scheme/unresolvable/private_ip/network_error/too_large family. +type Reason string + +const ( + ReasonInvalidURL Reason = "invalid_url" + ReasonScheme Reason = "scheme" + ReasonUnresolvable Reason = "unresolvable" + ReasonPrivateIP Reason = "private_ip" + ReasonNetworkError Reason = "network_error" + ReasonTooLarge Reason = "too_large" +) + +// Policy is supplied per call. +type Policy struct { + Mode Mode + AllowHosts []string // exact or dotted-suffix match; used in AllowHostsMode + MaxBytes int64 // 0 means use the config/framework default, never unlimited + Timeout time.Duration + + // tlsConfig, if set, is Transport.TLSClientConfig so tests can trust an + // httptest certificate. Production callers leave it nil. + tlsConfig *tls.Config + // skipReservedCheck disables the dial-time private-IP block so tests can + // exercise a real httptest listener on 127.0.0.1. Production callers + // leave it false. + skipReservedCheck bool +} + +// Error carries the typed Reason plus the underlying error for logging. +type Error struct { + Reason Reason + Err error +} + +func (e *Error) Error() string { + if e == nil { + return "fetchguard: error" + } + if e.Err != nil { + return fmt.Sprintf("fetchguard: %s: %v", e.Reason, e.Err) + } + return fmt.Sprintf("fetchguard: %s", e.Reason) +} + +func (e *Error) Unwrap() error { + if e == nil { + return nil + } + return e.Err +} + +// Defaults are the framework fallback: 10 MiB, 10s, matching PHP. +func Defaults() (maxBytes int64, timeout time.Duration) { + return 10 * 1024 * 1024, 10 * time.Second +} + +// DefaultsFromConfig reads http.fetch.max_bytes / http.fetch.timeout_seconds +// from cfg, falling back to Defaults() for absent keys. An explicitly +// configured zero or negative value is an error (D-14). +func DefaultsFromConfig(cfg *compass.Config) (maxBytes int64, timeout time.Duration, err error) { + maxBytes, timeout = Defaults() + if cfg == nil { + return maxBytes, timeout, nil + } + if v, ok := cfg.Lookup("http.fetch.max_bytes"); ok { + n, err := configInt64("http.fetch.max_bytes", v) + if err != nil { + return 0, 0, err + } + if n <= 0 { + return 0, 0, fmt.Errorf("fetchguard: http.fetch.max_bytes must be positive, got %d", n) + } + maxBytes = n + } + if v, ok := cfg.Lookup("http.fetch.timeout_seconds"); ok { + n, err := configInt64("http.fetch.timeout_seconds", v) + if err != nil { + return 0, 0, err + } + if n <= 0 { + return 0, 0, fmt.Errorf("fetchguard: http.fetch.timeout_seconds must be positive, got %d", n) + } + timeout = time.Duration(n) * time.Second + } + return maxBytes, timeout, nil +} + +func configInt64(path string, v any) (int64, error) { + switch n := v.(type) { + case int: + return int64(n), nil + case int8: + return int64(n), nil + case int16: + return int64(n), nil + case int32: + return int64(n), nil + case int64: + return n, nil + case uint: + return int64(n), nil + case uint8: + return int64(n), nil + case uint16: + return int64(n), nil + case uint32: + return int64(n), nil + case uint64: + if n > math.MaxInt64 { + return 0, fmt.Errorf("fetchguard: %s overflows int64", path) + } + return int64(n), nil + case float64: + if math.Trunc(n) != n { + return 0, fmt.Errorf("fetchguard: %s is not an integer", path) + } + if n > math.MaxInt64 || n < math.MinInt64 { + return 0, fmt.Errorf("fetchguard: %s overflows int64", path) + } + return int64(n), nil + case float32: + return configInt64(path, float64(n)) + default: + return 0, fmt.Errorf("fetchguard: %s has unexpected type %T", path, v) + } +}