97 lines
2.9 KiB
Go
97 lines
2.9 KiB
Go
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,
|
||
}
|
||
}
|