Files
chorus/internal/platform/http/safehttp.go
T

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
}