Files
pixa/internal/httpfetcher/httpfetcher.go
T
clawbot c7173c47d8
check / check (push) Waiting to run
Return and pass on request IDs, and give /v1/e/ ETag, 304 and HEAD (closes #84)
Every response carries X-Request-Id, the upstream fetch sends it, and
the "upstream fetched", "image converted" and "image served" lines log
it as request_id. pixa's own RequestID middleware keeps a request's own
ID only when it is at most 64 letters, digits, '-', '_' or '.', and
otherwise makes a random one with crypto/rand, so nothing a client
chooses freely and nothing about the host reaches upstream. /v1/e/ now
sets ETag, answers a matching If-None-Match with 304 and is routed for
HEAD, through notModified, which both image handlers call. No Vary is
added: go-chi/cors already sends Vary: Origin.

Model: opus-5-5
2026-10-04 12:41:54 +02:00

687 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
}
// 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()
}
// Create transport with SSRF-safe dialer. The dialer re-resolves and
// re-checks at connect time (closing the DNS-rebinding window) against
// both the built-in ranges and the operator-supplied blocklist.
transport := &http.Transport{
DialContext: func(ctx context.Context, network, addr string) (net.Conn, error) {
return dialSSRFSafe(ctx, network, addr, config.BlockedNetworks)
},
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
}