- 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
176 lines
4.4 KiB
Go
176 lines
4.4 KiB
Go
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}
|
|
}
|