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