285 lines
8.0 KiB
Go
285 lines
8.0 KiB
Go
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 "<invalid-url>"
|
|
}
|
|
return strings.ToLower(target.Scheme) + "://" + target.Host
|
|
}
|