httpfetcher.Config gets an optional DialContext. When it is set, New connects with it in place of the dialer that refuses internal addresses; the URL check and the redirect check still run. Nothing in the config file or the environment sets it, and pixa builds its fetcher without it, so production connects exactly as before. It lets a test outside this package send a public-looking address to a local test server. New tests check that a fetcher built without it refuses to connect to a local server, and that one built with it connects through it while a loopback URL and a redirect to a link-local address are still refused. Model: opus-5-5
698 lines
18 KiB
Go
698 lines
18 KiB
Go
// Package httpfetcher fetches content from upstream HTTP origins with SSRF
|
|
// protection, connection limits per host and for all hosts together, and
|
|
// content-type validation.
|
|
package httpfetcher
|
|
|
|
import (
|
|
"context"
|
|
"crypto/tls"
|
|
"errors"
|
|
"fmt"
|
|
"io"
|
|
"net"
|
|
"net/http"
|
|
"net/http/httptrace"
|
|
"net/netip"
|
|
neturl "net/url"
|
|
"slices"
|
|
"strings"
|
|
"sync"
|
|
"time"
|
|
|
|
"github.com/go-chi/chi/v5/middleware"
|
|
)
|
|
|
|
// Fetcher configuration constants.
|
|
const (
|
|
DefaultFetchTimeout = 30 * time.Second
|
|
DefaultMaxResponseSize = 50 << 20 // 50MB
|
|
DefaultTLSTimeout = 10 * time.Second
|
|
DefaultMaxIdleConns = 100
|
|
DefaultIdleConnTimeout = 90 * time.Second
|
|
DefaultMaxRedirects = 10
|
|
DefaultMaxConnectionsPerHost = 20
|
|
DefaultMaxConnections = 64
|
|
)
|
|
|
|
// ConnectionWaitTimeout is how long Fetch waits for a free connection when
|
|
// MaxConnections fetches are already in progress.
|
|
const ConnectionWaitTimeout = 10 * time.Second
|
|
|
|
// MIME content types.
|
|
const (
|
|
contentTypeJPEG = "image/jpeg"
|
|
contentTypePNG = "image/png"
|
|
contentTypeGIF = "image/gif"
|
|
contentTypeWebP = "image/webp"
|
|
contentTypeAVIF = "image/avif"
|
|
contentTypeSVG = "image/svg+xml"
|
|
contentTypeOctetStream = "application/octet-stream"
|
|
)
|
|
|
|
// Loopback addresses blocked by SSRF protection.
|
|
const (
|
|
localhostIPv4 = "127.0.0.1"
|
|
localhostIPv6 = "::1"
|
|
)
|
|
|
|
// builtinBlockedPrefixes are internal or special-use ranges that Go's
|
|
// net.IP predicates (IsPrivate, IsLinkLocalUnicast, and the like) do not
|
|
// already cover. They are always blocked, in addition to any
|
|
// operator-supplied networks. IPv4-mapped IPv6 addresses are unmapped
|
|
// before matching, so these IPv4 ranges are caught in both forms.
|
|
//
|
|
//nolint:gochecknoglobals // immutable built-in blocklist
|
|
var builtinBlockedPrefixes = []netip.Prefix{
|
|
netip.MustParsePrefix("100.64.0.0/10"), // RFC 6598 CGNAT / carrier-grade NAT
|
|
netip.MustParsePrefix("192.0.0.0/24"), // RFC 6890 IETF protocol assignments
|
|
netip.MustParsePrefix("198.18.0.0/15"), // RFC 2544 benchmarking range
|
|
netip.MustParsePrefix("64:ff9b::/96"), // RFC 6052 NAT64 (maps onto IPv4)
|
|
}
|
|
|
|
// Fetcher errors.
|
|
var (
|
|
ErrSSRFBlocked = errors.New("request blocked: private or internal IP")
|
|
ErrInvalidHost = errors.New("invalid or unresolvable host")
|
|
ErrUnsupportedScheme = errors.New("only HTTPS is supported")
|
|
ErrResponseTooLarge = errors.New("response exceeds maximum size")
|
|
ErrInvalidContentType = errors.New("invalid or unsupported content type")
|
|
ErrUpstreamError = errors.New("upstream server error")
|
|
ErrUpstreamTimeout = errors.New("upstream request timeout")
|
|
ErrTooManyConnections = errors.New("too many concurrent upstream connections")
|
|
)
|
|
|
|
// Internal fetcher errors.
|
|
var (
|
|
errTooManyRedirects = errors.New("too many redirects")
|
|
errConnectFailed = errors.New("failed to connect")
|
|
)
|
|
|
|
// Fetcher retrieves content from upstream origins.
|
|
type Fetcher interface {
|
|
// Fetch retrieves content from the given URL.
|
|
Fetch(ctx context.Context, url string) (*FetchResult, error)
|
|
}
|
|
|
|
// FetchResult contains the result of fetching from upstream.
|
|
type FetchResult struct {
|
|
// Content is the raw image data.
|
|
Content io.ReadCloser
|
|
// ContentLength is the size in bytes (-1 if unknown).
|
|
ContentLength int64
|
|
// ContentType is the MIME type from upstream.
|
|
ContentType string
|
|
// Headers contains all response headers from upstream.
|
|
Headers map[string][]string
|
|
// StatusCode is the HTTP status code from upstream.
|
|
StatusCode int
|
|
// FetchDurationMs is how long the fetch took in milliseconds.
|
|
FetchDurationMs int64
|
|
// RemoteAddr is the IP:port of the upstream server.
|
|
RemoteAddr string
|
|
// HTTPVersion is the protocol version (e.g., "1.1", "2.0").
|
|
HTTPVersion string
|
|
// TLSVersion is the TLS protocol version (e.g., "TLS 1.3").
|
|
TLSVersion string
|
|
// TLSCipherSuite is the negotiated cipher suite name.
|
|
TLSCipherSuite string
|
|
}
|
|
|
|
// Config holds configuration for the upstream fetcher.
|
|
type Config struct {
|
|
// Timeout for upstream requests.
|
|
Timeout time.Duration
|
|
// MaxResponseSize is the maximum allowed response body size.
|
|
MaxResponseSize int64
|
|
// UserAgent to send to upstream servers.
|
|
UserAgent string
|
|
// AllowedContentTypes is an allow list of MIME types to accept.
|
|
AllowedContentTypes []string
|
|
// AllowHTTP allows non-TLS connections (for testing only).
|
|
AllowHTTP bool
|
|
// MaxConnectionsPerHost limits concurrent connections to each upstream host.
|
|
MaxConnectionsPerHost int
|
|
// MaxConnections limits concurrent connections to all upstream hosts
|
|
// together.
|
|
MaxConnections int
|
|
// BlockedNetworks are operator-supplied CIDR ranges refused by the
|
|
// dialer, in addition to the always-enforced built-in ranges.
|
|
BlockedNetworks []netip.Prefix
|
|
// DialContext, when set, makes the fetcher's connections in place of
|
|
// the dialer that refuses internal addresses; the URL and redirect
|
|
// checks still run. Only tests set it, to reach a local server; the
|
|
// config file and the environment cannot.
|
|
DialContext func(ctx context.Context, network, addr string) (net.Conn, error)
|
|
}
|
|
|
|
// DefaultConfig returns a Config with sensible defaults.
|
|
func DefaultConfig() *Config {
|
|
return &Config{
|
|
Timeout: DefaultFetchTimeout,
|
|
MaxResponseSize: DefaultMaxResponseSize,
|
|
UserAgent: "pixa/1.0",
|
|
AllowedContentTypes: []string{
|
|
contentTypeJPEG,
|
|
contentTypePNG,
|
|
contentTypeGIF,
|
|
contentTypeWebP,
|
|
contentTypeAVIF,
|
|
contentTypeSVG,
|
|
},
|
|
AllowHTTP: false,
|
|
MaxConnectionsPerHost: DefaultMaxConnectionsPerHost,
|
|
MaxConnections: DefaultMaxConnections,
|
|
}
|
|
}
|
|
|
|
// HTTPFetcher implements Fetcher with SSRF protection and connection limits
|
|
// per host and for all hosts together.
|
|
type HTTPFetcher struct {
|
|
client *http.Client
|
|
config *Config
|
|
// hostSems holds the semaphore of each host with a fetch holding or
|
|
// waiting for one of its slots; the entry is removed when the host's
|
|
// last such fetch gives its slot back or stops waiting.
|
|
hostSems map[string]*hostSemaphore
|
|
hostSemMu sync.Mutex // protects hostSems and each entry's count
|
|
// allHostsSemaphore has one slot per connection allowed to all hosts
|
|
// together (config.MaxConnections).
|
|
allHostsSemaphore chan struct{}
|
|
// connectionWaitTimeout is ConnectionWaitTimeout; tests shorten it.
|
|
connectionWaitTimeout time.Duration
|
|
}
|
|
|
|
// hostSemaphore is one host's connection slots
|
|
// (config.MaxConnectionsPerHost) and the number of fetches holding or
|
|
// waiting for one of them.
|
|
type hostSemaphore struct {
|
|
slots chan struct{}
|
|
count int
|
|
}
|
|
|
|
// New creates a new HTTPFetcher with SSRF protection.
|
|
func New(config *Config) *HTTPFetcher {
|
|
if config == nil {
|
|
config = DefaultConfig()
|
|
}
|
|
|
|
// Unless config.DialContext replaces it, the transport connects with
|
|
// the SSRF-safe dialer, which re-resolves and re-checks at connect time
|
|
// (closing the DNS-rebinding window) against both the built-in ranges
|
|
// and the operator-supplied blocklist.
|
|
dialContext := config.DialContext
|
|
if dialContext == nil {
|
|
dialContext = func(ctx context.Context, network, addr string) (net.Conn, error) {
|
|
return dialSSRFSafe(ctx, network, addr, config.BlockedNetworks)
|
|
}
|
|
}
|
|
|
|
transport := &http.Transport{
|
|
DialContext: dialContext,
|
|
TLSHandshakeTimeout: DefaultTLSTimeout,
|
|
MaxIdleConns: DefaultMaxIdleConns,
|
|
IdleConnTimeout: DefaultIdleConnTimeout,
|
|
}
|
|
|
|
client := &http.Client{
|
|
Transport: transport,
|
|
Timeout: config.Timeout,
|
|
// Don't follow redirects automatically - we need to validate each hop
|
|
CheckRedirect: func(req *http.Request, via []*http.Request) error {
|
|
if len(via) >= DefaultMaxRedirects {
|
|
return errTooManyRedirects
|
|
}
|
|
|
|
// Validate the redirect target
|
|
err := validateURL(req.Context(), req.URL.String(), config.AllowHTTP)
|
|
if err != nil {
|
|
return fmt.Errorf("redirect blocked: %w", err)
|
|
}
|
|
|
|
return nil
|
|
},
|
|
}
|
|
|
|
return &HTTPFetcher{
|
|
client: client,
|
|
config: config,
|
|
hostSems: make(map[string]*hostSemaphore),
|
|
allHostsSemaphore: make(chan struct{}, config.MaxConnections),
|
|
connectionWaitTimeout: ConnectionWaitTimeout,
|
|
}
|
|
}
|
|
|
|
// Fetch retrieves content from the given URL with SSRF protection. When
|
|
// MaxConnections fetches are already in progress, it waits up to
|
|
// ConnectionWaitTimeout for one to finish, then fails with
|
|
// ErrTooManyConnections.
|
|
func (f *HTTPFetcher) Fetch(ctx context.Context, url string) (*FetchResult, error) {
|
|
// Validate URL before making request
|
|
err := validateURL(ctx, url, f.config.AllowHTTP)
|
|
if err != nil {
|
|
return nil, err
|
|
}
|
|
|
|
release, err := f.acquireConnection(ctx, extractHost(url))
|
|
if err != nil {
|
|
return nil, err
|
|
}
|
|
|
|
// If we fail before returning a result, release the connection
|
|
success := false
|
|
|
|
defer func() {
|
|
if !success {
|
|
release()
|
|
}
|
|
}()
|
|
|
|
parsedURL, err := neturl.Parse(url)
|
|
if err != nil {
|
|
return nil, fmt.Errorf("failed to parse URL: %w", err)
|
|
}
|
|
|
|
req := &http.Request{
|
|
Method: http.MethodGet,
|
|
URL: parsedURL,
|
|
Header: make(http.Header),
|
|
}
|
|
|
|
req.Header.Set("User-Agent", f.config.UserAgent)
|
|
req.Header.Set("Accept", strings.Join(f.config.AllowedContentTypes, ", "))
|
|
|
|
// The ID of the request this fetch serves, so the fetch can be found in
|
|
// the upstream host's logs
|
|
requestID := middleware.GetReqID(ctx)
|
|
if requestID != "" {
|
|
req.Header.Set(middleware.RequestIDHeader, requestID)
|
|
}
|
|
|
|
// Use httptrace to capture connection details
|
|
var remoteAddr string
|
|
|
|
trace := &httptrace.ClientTrace{
|
|
GotConn: func(info httptrace.GotConnInfo) {
|
|
if info.Conn != nil {
|
|
remoteAddr = info.Conn.RemoteAddr().String()
|
|
}
|
|
},
|
|
}
|
|
req = req.WithContext(httptrace.WithClientTrace(ctx, trace))
|
|
|
|
startTime := time.Now()
|
|
|
|
resp, err := f.client.Do(req)
|
|
|
|
fetchDuration := time.Since(startTime)
|
|
|
|
if err != nil {
|
|
if errors.Is(err, context.DeadlineExceeded) {
|
|
return nil, ErrUpstreamTimeout
|
|
}
|
|
|
|
return nil, fmt.Errorf("upstream request failed: %w", err)
|
|
}
|
|
|
|
result, err := f.buildResult(resp, remoteAddr, fetchDuration, release)
|
|
if err != nil {
|
|
return nil, err
|
|
}
|
|
|
|
// Mark success so defer doesn't release the connection; closing the
|
|
// result's Content does
|
|
success = true
|
|
|
|
return result, nil
|
|
}
|
|
|
|
// acquireConnection takes a slot for host, then one of the slots shared by
|
|
// all hosts, and returns the func that gives both back. The host's slot
|
|
// comes first, so fetches queued for one busy host hold no shared slot.
|
|
// Only the wait for a shared slot is bounded: after connectionWaitTimeout
|
|
// it fails with ErrTooManyConnections.
|
|
func (f *HTTPFetcher) acquireConnection(
|
|
ctx context.Context, host string,
|
|
) (func(), error) {
|
|
hostSem := f.getHostSemaphore(host)
|
|
|
|
select {
|
|
case hostSem <- struct{}{}:
|
|
case <-ctx.Done():
|
|
f.putHostSemaphore(host)
|
|
|
|
return nil, ctx.Err()
|
|
}
|
|
|
|
select {
|
|
case f.allHostsSemaphore <- struct{}{}:
|
|
case <-time.After(f.connectionWaitTimeout):
|
|
<-hostSem
|
|
f.putHostSemaphore(host)
|
|
|
|
return nil, ErrTooManyConnections
|
|
case <-ctx.Done():
|
|
<-hostSem
|
|
f.putHostSemaphore(host)
|
|
|
|
return nil, ctx.Err()
|
|
}
|
|
|
|
return func() {
|
|
<-hostSem
|
|
f.putHostSemaphore(host)
|
|
<-f.allHostsSemaphore
|
|
}, nil
|
|
}
|
|
|
|
// getHostSemaphore returns the semaphore for a host, creating it if
|
|
// necessary, and counts the caller among the fetches using it. The caller
|
|
// calls putHostSemaphore once it holds no slot and waits for none.
|
|
func (f *HTTPFetcher) getHostSemaphore(host string) chan struct{} {
|
|
f.hostSemMu.Lock()
|
|
defer f.hostSemMu.Unlock()
|
|
|
|
sem, ok := f.hostSems[host]
|
|
if !ok {
|
|
sem = &hostSemaphore{
|
|
slots: make(chan struct{}, f.config.MaxConnectionsPerHost),
|
|
}
|
|
f.hostSems[host] = sem
|
|
}
|
|
|
|
sem.count++
|
|
|
|
return sem.slots
|
|
}
|
|
|
|
// putHostSemaphore stops counting the caller among the fetches using the
|
|
// host's semaphore, and removes the semaphore when no fetch uses it.
|
|
func (f *HTTPFetcher) putHostSemaphore(host string) {
|
|
f.hostSemMu.Lock()
|
|
defer f.hostSemMu.Unlock()
|
|
|
|
sem := f.hostSems[host]
|
|
|
|
sem.count--
|
|
if sem.count == 0 {
|
|
delete(f.hostSems, host)
|
|
}
|
|
}
|
|
|
|
// buildResult validates the upstream response and assembles a FetchResult
|
|
// whose Content calls release when closed.
|
|
func (f *HTTPFetcher) buildResult(
|
|
resp *http.Response,
|
|
remoteAddr string,
|
|
fetchDuration time.Duration,
|
|
release func(),
|
|
) (*FetchResult, error) {
|
|
// Extract HTTP version (strip "HTTP/" prefix)
|
|
httpVersion := strings.TrimPrefix(resp.Proto, "HTTP/")
|
|
|
|
// Extract TLS info if available
|
|
var tlsVersion, tlsCipherSuite string
|
|
|
|
if resp.TLS != nil {
|
|
tlsVersion = tls.VersionName(resp.TLS.Version)
|
|
tlsCipherSuite = tls.CipherSuiteName(resp.TLS.CipherSuite)
|
|
}
|
|
|
|
// Check status code
|
|
if resp.StatusCode < 200 || resp.StatusCode >= 300 {
|
|
_ = resp.Body.Close()
|
|
|
|
return nil, fmt.Errorf("%w: status %d", ErrUpstreamError, resp.StatusCode)
|
|
}
|
|
|
|
// Validate content type
|
|
contentType := resp.Header.Get("Content-Type")
|
|
if !f.isAllowedContentType(contentType) {
|
|
_ = resp.Body.Close()
|
|
|
|
return nil, fmt.Errorf("%w: %s", ErrInvalidContentType, contentType)
|
|
}
|
|
|
|
// Wrap body with size limiter and semaphore releaser
|
|
limitedBody := &limitedReader{
|
|
reader: resp.Body,
|
|
remaining: f.config.MaxResponseSize,
|
|
}
|
|
|
|
return &FetchResult{
|
|
Content: &semaphoreReleasingReadCloser{limitedBody, resp.Body, release},
|
|
ContentLength: resp.ContentLength,
|
|
ContentType: contentType,
|
|
Headers: resp.Header,
|
|
StatusCode: resp.StatusCode,
|
|
FetchDurationMs: fetchDuration.Milliseconds(),
|
|
RemoteAddr: remoteAddr,
|
|
HTTPVersion: httpVersion,
|
|
TLSVersion: tlsVersion,
|
|
TLSCipherSuite: tlsCipherSuite,
|
|
}, nil
|
|
}
|
|
|
|
// isAllowedContentType checks if the content type is in the allow list.
|
|
func (f *HTTPFetcher) isAllowedContentType(contentType string) bool {
|
|
// Extract the MIME type without parameters
|
|
mediaType := strings.TrimSpace(strings.Split(contentType, ";")[0])
|
|
|
|
for _, allowed := range f.config.AllowedContentTypes {
|
|
if strings.EqualFold(mediaType, allowed) {
|
|
return true
|
|
}
|
|
}
|
|
|
|
return false
|
|
}
|
|
|
|
// validateURL checks if a URL is safe to fetch (not internal/private).
|
|
func validateURL(ctx context.Context, rawURL string, allowHTTP bool) error {
|
|
if !allowHTTP && !strings.HasPrefix(rawURL, "https://") {
|
|
return ErrUnsupportedScheme
|
|
}
|
|
|
|
// Parse to extract host
|
|
host := extractHost(rawURL)
|
|
if host == "" {
|
|
return ErrInvalidHost
|
|
}
|
|
|
|
// Remove port if present
|
|
h, _, err := net.SplitHostPort(host)
|
|
if err == nil {
|
|
host = h
|
|
}
|
|
|
|
// Block obvious localhost patterns
|
|
if isLocalhost(host) {
|
|
return ErrSSRFBlocked
|
|
}
|
|
|
|
// Resolve the host to check IP addresses
|
|
addrs, err := net.DefaultResolver.LookupIPAddr(ctx, host)
|
|
if err != nil {
|
|
return fmt.Errorf("%w: %s", ErrInvalidHost, host)
|
|
}
|
|
|
|
private := slices.ContainsFunc(addrs, func(addr net.IPAddr) bool {
|
|
return isPrivateIP(addr.IP)
|
|
})
|
|
if private {
|
|
return ErrSSRFBlocked
|
|
}
|
|
|
|
return nil
|
|
}
|
|
|
|
// extractHost extracts the host from a URL string.
|
|
func extractHost(rawURL string) string {
|
|
// Simple extraction without full URL parsing
|
|
url := rawURL
|
|
if idx := strings.Index(url, "://"); idx != -1 {
|
|
url = url[idx+3:]
|
|
}
|
|
|
|
if idx := strings.Index(url, "/"); idx != -1 {
|
|
url = url[:idx]
|
|
}
|
|
|
|
if idx := strings.Index(url, "?"); idx != -1 {
|
|
url = url[:idx]
|
|
}
|
|
|
|
return url
|
|
}
|
|
|
|
// isLocalhost checks if the host is localhost.
|
|
func isLocalhost(host string) bool {
|
|
host = strings.ToLower(host)
|
|
|
|
return host == "localhost" ||
|
|
host == localhostIPv4 ||
|
|
host == localhostIPv6 ||
|
|
host == "[::1]" ||
|
|
strings.HasSuffix(host, ".localhost") ||
|
|
strings.HasSuffix(host, ".local")
|
|
}
|
|
|
|
// isPrivateIP checks if an IP is private, loopback, or otherwise internal.
|
|
func isPrivateIP(ip net.IP) bool {
|
|
if ip == nil {
|
|
return true
|
|
}
|
|
|
|
// Check for loopback
|
|
if ip.IsLoopback() {
|
|
return true
|
|
}
|
|
|
|
// Check for private ranges
|
|
if ip.IsPrivate() {
|
|
return true
|
|
}
|
|
|
|
// Check for link-local
|
|
if ip.IsLinkLocalUnicast() || ip.IsLinkLocalMulticast() {
|
|
return true
|
|
}
|
|
|
|
// Check for unspecified (0.0.0.0 or ::)
|
|
if ip.IsUnspecified() {
|
|
return true
|
|
}
|
|
|
|
// Check for multicast
|
|
if ip.IsMulticast() {
|
|
return true
|
|
}
|
|
|
|
// Additional checks for IPv4
|
|
if ip4 := ip.To4(); ip4 != nil {
|
|
// 169.254.0.0/16 - Link local
|
|
if ip4[0] == 169 && ip4[1] == 254 {
|
|
return true
|
|
}
|
|
// 0.0.0.0/8 - Current network
|
|
if ip4[0] == 0 {
|
|
return true
|
|
}
|
|
}
|
|
|
|
// Special-use ranges the net.IP predicates above do not cover.
|
|
addr, ok := netip.AddrFromSlice(ip)
|
|
if !ok {
|
|
return true
|
|
}
|
|
|
|
addr = addr.Unmap()
|
|
|
|
return slices.ContainsFunc(builtinBlockedPrefixes, func(prefix netip.Prefix) bool {
|
|
return prefix.Contains(addr)
|
|
})
|
|
}
|
|
|
|
// isBlockedIP reports whether ip is refused, either by the built-in
|
|
// internal-range check or by one of the operator-supplied prefixes.
|
|
func isBlockedIP(ip net.IP, blocked []netip.Prefix) bool {
|
|
if isPrivateIP(ip) {
|
|
return true
|
|
}
|
|
|
|
addr, ok := netip.AddrFromSlice(ip)
|
|
if !ok {
|
|
return true
|
|
}
|
|
|
|
addr = addr.Unmap()
|
|
|
|
return slices.ContainsFunc(blocked, func(prefix netip.Prefix) bool {
|
|
return prefix.Contains(addr)
|
|
})
|
|
}
|
|
|
|
// ssrfSafeDialer validates IP addresses against the built-in blocked ranges
|
|
// before connecting. New wraps dialSSRFSafe with the operator-supplied
|
|
// blocklist; this entry point enforces the built-in ranges alone.
|
|
func ssrfSafeDialer(ctx context.Context, network, addr string) (net.Conn, error) {
|
|
return dialSSRFSafe(ctx, network, addr, nil)
|
|
}
|
|
|
|
// dialSSRFSafe re-resolves addr and refuses to connect to any built-in
|
|
// internal range or operator-supplied blocked prefix, closing the
|
|
// DNS-rebinding window at connect time.
|
|
func dialSSRFSafe(
|
|
ctx context.Context,
|
|
network, addr string,
|
|
blocked []netip.Prefix,
|
|
) (net.Conn, error) {
|
|
host, port, err := net.SplitHostPort(addr)
|
|
if err != nil {
|
|
return nil, err
|
|
}
|
|
|
|
// Resolve the address
|
|
ips, err := net.DefaultResolver.LookupIP(ctx, "ip", host)
|
|
if err != nil {
|
|
return nil, fmt.Errorf("%w: %s", ErrInvalidHost, host)
|
|
}
|
|
|
|
// Check all resolved IPs
|
|
for _, ip := range ips {
|
|
if isBlockedIP(ip, blocked) {
|
|
return nil, ErrSSRFBlocked
|
|
}
|
|
}
|
|
|
|
// Connect using the first valid IP
|
|
var dialer net.Dialer
|
|
|
|
for _, ip := range ips {
|
|
addr := net.JoinHostPort(ip.String(), port)
|
|
|
|
conn, err := dialer.DialContext(ctx, network, addr)
|
|
if err == nil {
|
|
return conn, nil
|
|
}
|
|
}
|
|
|
|
return nil, fmt.Errorf("%w to %s", errConnectFailed, host)
|
|
}
|
|
|
|
// limitedReader wraps a reader and limits the number of bytes read.
|
|
type limitedReader struct {
|
|
reader io.Reader
|
|
remaining int64
|
|
}
|
|
|
|
func (r *limitedReader) Read(p []byte) (int, error) {
|
|
if r.remaining <= 0 {
|
|
return 0, ErrResponseTooLarge
|
|
}
|
|
|
|
if int64(len(p)) > r.remaining {
|
|
p = p[:r.remaining]
|
|
}
|
|
|
|
n, err := r.reader.Read(p)
|
|
r.remaining -= int64(n)
|
|
|
|
return n, err
|
|
}
|
|
|
|
// semaphoreReleasingReadCloser releases the fetch's connection slots when
|
|
// closed.
|
|
type semaphoreReleasingReadCloser struct {
|
|
*limitedReader
|
|
|
|
closer io.Closer
|
|
release func()
|
|
}
|
|
|
|
func (r *semaphoreReleasingReadCloser) Close() error {
|
|
err := r.closer.Close()
|
|
r.release()
|
|
|
|
return err
|
|
}
|