package safehttp import ( "context" "errors" "fmt" "io" "net" "net/http" "net/netip" "net/url" "strconv" "strings" "time" ) var ( ErrInvalidTarget = errors.New("outbound target is invalid") ErrBlockedAddress = errors.New("outbound target address is blocked") ErrTooManyRedirects = errors.New("outbound redirect limit exceeded") ErrCrossOriginAuth = errors.New("authenticated cross-origin redirect is blocked") ErrResponseTooLarge = errors.New("outbound response exceeds size limit") ErrRequestFailed = errors.New("safe outbound request failed") ) type Resolver interface { LookupIPAddr(ctx context.Context, host string) ([]net.IPAddr, error) } type DialContextFunc func(ctx context.Context, network, address string) (net.Conn, error) type Config struct { Resolver Resolver DialContext DialContextFunc Timeout time.Duration MaxRedirects int AllowedPorts []uint16 } type Client struct { httpClient *http.Client validator *validator } type Response struct { StatusCode int ContentType string Body []byte } type Request struct { Method string URL string Header http.Header Body io.Reader MaxBytes int64 } func New(config Config) (*Client, error) { if config.Resolver == nil { config.Resolver = net.DefaultResolver } if config.DialContext == nil { dialer := &net.Dialer{Timeout: 10 * time.Second, KeepAlive: 30 * time.Second} config.DialContext = dialer.DialContext } if config.Timeout <= 0 { return nil, fmt.Errorf("safe HTTP timeout must be positive") } if config.MaxRedirects < 0 { return nil, fmt.Errorf("safe HTTP redirect limit cannot be negative") } if len(config.AllowedPorts) == 0 { config.AllowedPorts = []uint16{80, 443} } ports := make(map[string]struct{}, len(config.AllowedPorts)) for _, port := range config.AllowedPorts { if port == 0 { return nil, fmt.Errorf("safe HTTP allowed port cannot be zero") } ports[strconv.Itoa(int(port))] = struct{}{} } v := &validator{resolver: config.Resolver, allowedPorts: ports} safeDialer := &dialer{validator: v, dial: config.DialContext} transport := &http.Transport{ Proxy: nil, DialContext: safeDialer.DialContext, ForceAttemptHTTP2: true, MaxIdleConns: 20, IdleConnTimeout: 30 * time.Second, TLSHandshakeTimeout: 10 * time.Second, ResponseHeaderTimeout: config.Timeout, } client := &http.Client{Transport: transport, Timeout: config.Timeout} client.CheckRedirect = func(request *http.Request, via []*http.Request) error { if len(via) > config.MaxRedirects { return ErrTooManyRedirects } if len(via) > 0 && via[0].Header.Get("Authorization") != "" && !sameOrigin(via[0].URL, request.URL) { return ErrCrossOriginAuth } return v.validateURL(request.Context(), request.URL, true) } return &Client{httpClient: client, validator: v}, nil } func sameOrigin(first, second *url.URL) bool { if first == nil || second == nil { return false } return strings.EqualFold(first.Scheme, second.Scheme) && strings.EqualFold(first.Host, second.Host) } func (c *Client) Fetch(ctx context.Context, rawURL string, maxBytes int64) (Response, error) { return c.Do(ctx, Request{Method: http.MethodGet, URL: rawURL, MaxBytes: maxBytes}) } func (c *Client) Do(ctx context.Context, input Request) (Response, error) { if input.MaxBytes <= 0 { return Response{}, fmt.Errorf("maximum response bytes must be positive") } method := strings.ToUpper(strings.TrimSpace(input.Method)) if method != http.MethodGet && method != http.MethodPost { return Response{}, fmt.Errorf("safe HTTP method is not allowed") } target, err := url.Parse(input.URL) if err != nil { return Response{}, ErrInvalidTarget } if err := c.validator.validateURL(ctx, target, true); err != nil { return Response{}, err } request, err := http.NewRequestWithContext(ctx, method, target.String(), input.Body) if err != nil { return Response{}, ErrInvalidTarget } request.Header = input.Header.Clone() response, err := c.httpClient.Do(request) if err != nil { var urlError *url.Error if errors.As(err, &urlError) && urlError.Err != nil { return Response{}, fmt.Errorf("%w: %w", ErrRequestFailed, urlError.Err) } return Response{}, ErrRequestFailed } defer response.Body.Close() if response.ContentLength > input.MaxBytes { return Response{}, ErrResponseTooLarge } body, err := io.ReadAll(io.LimitReader(response.Body, input.MaxBytes+1)) if err != nil { return Response{}, fmt.Errorf("read safe HTTP response: %w", err) } if int64(len(body)) > input.MaxBytes { return Response{}, ErrResponseTooLarge } return Response{StatusCode: response.StatusCode, ContentType: response.Header.Get("Content-Type"), Body: body}, nil } func (c *Client) CloseIdleConnections() { c.httpClient.CloseIdleConnections() } type validator struct { resolver Resolver allowedPorts map[string]struct{} } func (v *validator) validateURL(ctx context.Context, target *url.URL, resolve bool) error { if target == nil || (target.Scheme != "http" && target.Scheme != "https") || target.Hostname() == "" || target.User != nil { return ErrInvalidTarget } port := target.Port() if port == "" { if target.Scheme == "https" { port = "443" } else { port = "80" } } if _, ok := v.allowedPorts[port]; !ok { return ErrInvalidTarget } if resolve { _, err := v.resolvePublic(ctx, target.Hostname()) return err } return nil } func (v *validator) resolvePublic(ctx context.Context, host string) ([]net.IPAddr, error) { if parsed := net.ParseIP(host); parsed != nil { addresses := []net.IPAddr{{IP: parsed}} if err := validateAddresses(addresses); err != nil { return nil, err } return addresses, nil } addresses, err := v.resolver.LookupIPAddr(ctx, host) if err != nil || len(addresses) == 0 { return nil, ErrInvalidTarget } if err := validateAddresses(addresses); err != nil { return nil, err } return addresses, nil } type dialer struct { validator *validator dial DialContextFunc } func (d *dialer) DialContext(ctx context.Context, network, address string) (net.Conn, error) { host, port, err := net.SplitHostPort(address) if err != nil { return nil, ErrInvalidTarget } if _, ok := d.validator.allowedPorts[port]; !ok { return nil, ErrInvalidTarget } addresses, err := d.validator.resolvePublic(ctx, host) if err != nil { return nil, err } var lastErr error for _, resolved := range addresses { ip := resolved.IP.String() connection, dialErr := d.dial(ctx, network, net.JoinHostPort(ip, port)) if dialErr == nil { return connection, nil } lastErr = dialErr } if lastErr != nil { return nil, fmt.Errorf("safe outbound connection failed") } return nil, ErrInvalidTarget } func validateAddresses(addresses []net.IPAddr) error { for _, address := range addresses { parsed, ok := netip.AddrFromSlice(address.IP) if !ok || blockedIP(parsed.Unmap()) { return ErrBlockedAddress } } return nil } var blockedPrefixes = mustPrefixes( "0.0.0.0/8", "10.0.0.0/8", "100.64.0.0/10", "127.0.0.0/8", "169.254.0.0/16", "172.16.0.0/12", "192.0.0.0/24", "192.0.2.0/24", "192.88.99.0/24", "192.168.0.0/16", "198.18.0.0/15", "198.51.100.0/24", "203.0.113.0/24", "224.0.0.0/4", "240.0.0.0/4", "::/128", "::1/128", "64:ff9b::/96", "100::/64", "2001:db8::/32", "fc00::/7", "fe80::/10", "ff00::/8", ) func blockedIP(address netip.Addr) bool { if !address.IsValid() || !address.IsGlobalUnicast() || address.IsPrivate() || address.IsLoopback() || address.IsLinkLocalUnicast() || address.IsLinkLocalMulticast() || address.IsMulticast() || address.IsUnspecified() { return true } for _, prefix := range blockedPrefixes { if prefix.Contains(address) { return true } } return false } func mustPrefixes(values ...string) []netip.Prefix { prefixes := make([]netip.Prefix, 0, len(values)) for _, value := range values { prefixes = append(prefixes, netip.MustParsePrefix(value)) } return prefixes } func RedactedURL(rawURL string) string { target, err := url.Parse(rawURL) if err != nil || target.Hostname() == "" { return "" } return strings.ToLower(target.Scheme) + "://" + target.Host }