Files
2026-08-16 16:57:36 +02:00

97 lines
2.9 KiB
Go
Raw Permalink Blame History

This file contains ambiguous Unicode characters
This file contains Unicode characters that might be confused with other characters. If you think that this is intentional, you can safely ignore this warning. Use the Escape button to reveal them.
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 60180s).
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,
}
}