Files
summercms/modules/fetchguard/fetch.go
Jakub Zych 5e50b166ef refactor(10.2-01): nest framework packages under modules
- Move remaining beach packages and embedded admin assets\n- Rewrite framework, example, build, and gate paths
2026-09-28 02:21:02 +02:00

179 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/modules/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)
}
if addr.Zone() != "" {
return errPrivateIP
}
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}
}