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 } tr := SafeHTTPTransportPolicy(policy) // Match header wait to overall timeout. A fixed 30s ResponseHeaderTimeout // aborts OpenAI-compatible LLM calls that withhold headers until generation // finishes (reasoning models often take 60–180s). tr.ResponseHeaderTimeout = timeout return &http.Client{ Timeout: timeout, Transport: tr, } } // 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, } }