feat(06-04): implement dial-time SSRF-guarded Fetch
- AllowHosts exact/dotted-suffix and PublicOnly modes, https-only - Private/reserved IP check at net.Dialer.Control, no redirects - Streaming io.LimitReader byte cap and typed failure reasons
This commit is contained in:
175
fetchguard/fetch.go
Normal file
175
fetchguard/fetch.go
Normal file
@@ -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}
|
||||
}
|
||||
147
fetchguard/policy.go
Normal file
147
fetchguard/policy.go
Normal file
@@ -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)
|
||||
}
|
||||
}
|
||||
Reference in New Issue
Block a user