Files
descrybe/apps/api/internal/security/http_client.go
T

92 lines
2.7 KiB
Go
Raw Normal View History

package security
import (
"context"
"fmt"
"net"
"net/http"
"time"
)
const (
defaultDialTimeout = 10 * time.Second
defaultTLSHandshakeTimeout = 10 * time.Second
defaultResponseHeaderTimeout = 30 * time.Second
)
// SafeHTTPClient returns an HTTP client whose DialContext refuses private,
// link-local, CGNAT, and cloud-metadata addresses (SSRF). When allowLoopback
// is true, localhost is permitted (local Woo/Shopify mocks only).
func SafeHTTPClient(timeout time.Duration, allowLoopback bool) *http.Client {
return SafeHTTPClientPolicy(timeout, DialPolicy{AllowLoopback: allowLoopback})
}
// SafeHTTPClientPolicy is SafeHTTPClient with an explicit DialPolicy.
func SafeHTTPClientPolicy(timeout time.Duration, policy DialPolicy) *http.Client {
if timeout <= 0 {
timeout = 30 * time.Second
}
return &http.Client{
Timeout: timeout,
Transport: SafeHTTPTransportPolicy(policy),
}
}
// SafeHTTPTransport builds a transport with dial-time SSRF checks.
func SafeHTTPTransport(allowLoopback bool) *http.Transport {
return SafeHTTPTransportPolicy(DialPolicy{AllowLoopback: allowLoopback})
}
// SafeHTTPTransportPolicy builds a transport with dial-time SSRF checks per policy.
// Proxy is intentionally nil: HTTP(S)_PROXY would dial the proxy host and skip
// destination IP checks, defeating SSRF controls for user-influenced URLs.
func SafeHTTPTransportPolicy(policy DialPolicy) *http.Transport {
dialer := &net.Dialer{Timeout: defaultDialTimeout, KeepAlive: 30 * time.Second}
return &http.Transport{
Proxy: nil,
DialContext: func(ctx context.Context, network, addr string) (net.Conn, error) {
host, port, err := net.SplitHostPort(addr)
if err != nil {
return nil, err
}
if err := AssertHost(ctx, host, policy); err != nil {
return nil, fmt.Errorf("%w: %s", ErrBlockedHost, host)
}
ips, err := resolveHostIPs(ctx, host)
if err != nil {
return nil, err
}
var lastErr error
for _, ip := range ips {
if ip.IsLoopback() {
if !policy.AllowLoopback {
lastErr = ErrBlockedHost
continue
}
} else if isBlockedIP(ip) {
if !(policy.AllowPrivate && isPrivateLANIP(ip)) {
lastErr = ErrBlockedHost
continue
}
}
target := net.JoinHostPort(ip.String(), port)
conn, err := dialer.DialContext(ctx, network, target)
if err == nil {
return conn, nil
}
lastErr = err
}
if lastErr == nil {
lastErr = ErrBlockedHost
}
return nil, lastErr
},
ForceAttemptHTTP2: true,
MaxIdleConns: 10,
IdleConnTimeout: 90 * time.Second,
TLSHandshakeTimeout: defaultTLSHandshakeTimeout,
ExpectContinueTimeout: 1 * time.Second,
ResponseHeaderTimeout: defaultResponseHeaderTimeout,
}
}