250 lines
7.6 KiB
Go
250 lines
7.6 KiB
Go
package security
|
|
|
|
// Reference: P14 Deliverable #2 — SSRF Protection
|
|
// Prevents Server-Side Request Forgery attacks in outbound HTTP requests.
|
|
// Chatwoot uses lib/safe_fetch.rb for webhook URL validation; gochat needs equivalent.
|
|
|
|
import (
|
|
"context"
|
|
"fmt"
|
|
"net"
|
|
"net/http"
|
|
"net/url"
|
|
"strings"
|
|
"time"
|
|
)
|
|
|
|
// --- Security Audit Findings ---
|
|
//
|
|
// 1. CRITICAL: Webhook handler accepts arbitrary channel_type and inbox_id from URL params
|
|
// with no SSRF validation. When providers make outbound HTTP calls (e.g., Telegram API),
|
|
// an attacker controlling inbox config could redirect to internal services.
|
|
//
|
|
// 2. HIGH: No validation on URLs that providers might fetch. OAuth callback URLs,
|
|
// webhook verification URLs, and avatar URLs are all potential SSRF vectors.
|
|
//
|
|
// 3. MEDIUM: No DNS rebinding prevention. An attacker could register a domain that
|
|
// resolves to an internal IP after initial resolution.
|
|
//
|
|
// Chatwoot's safe_fetch.rb validates:
|
|
// - Resolves hostname and blocks private/reserved IPs
|
|
// - Blocks link-local, loopback, and multicast addresses
|
|
// - Uses custom DNS resolver to prevent rebinding
|
|
|
|
// SSRFConfig holds SSRF protection configuration.
|
|
type SSRFConfig struct {
|
|
AllowedDomains []string // retained trusted-domain catalog; never bypasses IP validation
|
|
BlockedCIDRs []string // IP ranges forbidden (private, loopback, etc.)
|
|
AllowedPorts []string // empty permits all ports; defaults restrict fetches to HTTP(S)
|
|
MaxRedirects int // limit HTTP redirect chains
|
|
RequireTLS bool // enforce HTTPS for certain operations
|
|
}
|
|
|
|
// DefaultSSRFConfig returns safe defaults matching Chatwoot's safe_fetch.rb.
|
|
func DefaultSSRFConfig() SSRFConfig {
|
|
return SSRFConfig{
|
|
AllowedDomains: []string{
|
|
"api.telegram.org",
|
|
"graph.facebook.com",
|
|
"api.instagram.com",
|
|
"business.facebook.com",
|
|
"web.whatsapp.com",
|
|
},
|
|
BlockedCIDRs: []string{
|
|
"10.0.0.0/8", // RFC 1918 private
|
|
"172.16.0.0/12", // RFC 1918 private
|
|
"192.168.0.0/16", // RFC 1918 private
|
|
"127.0.0.0/8", // Loopback
|
|
"0.0.0.0/8", // Current network
|
|
"100.64.0.0/10", // CGN
|
|
"169.254.0.0/16", // Link-local
|
|
"192.0.0.0/24", // IETF Protocol Assignments
|
|
"192.0.2.0/24", // TEST-NET-1
|
|
"198.51.100.0/24", // TEST-NET-2
|
|
"203.0.113.0/24", // TEST-NET-3
|
|
"198.18.0.0/15", // Benchmarking
|
|
"224.0.0.0/4", // Multicast
|
|
"240.0.0.0/4", // Reserved
|
|
"::1/128", // IPv6 loopback
|
|
"::/128", // IPv6 unspecified
|
|
"fc00::/7", // IPv6 unique local
|
|
"fe80::/10", // IPv6 link-local
|
|
"ff00::/8", // IPv6 multicast
|
|
"2001:db8::/32", // IPv6 documentation
|
|
},
|
|
AllowedPorts: []string{"80", "443"},
|
|
MaxRedirects: 3,
|
|
RequireTLS: false,
|
|
}
|
|
}
|
|
|
|
// SafeHTTPClient wraps http.Client with SSRF protection.
|
|
type SafeHTTPClient struct {
|
|
client *http.Client
|
|
cfg SSRFConfig
|
|
}
|
|
|
|
// NewSafeHTTPClient creates an HTTP client that blocks requests to private IPs.
|
|
func NewSafeHTTPClient(cfg SSRFConfig) *SafeHTTPClient {
|
|
dialer := &net.Dialer{
|
|
Timeout: 10 * time.Second,
|
|
KeepAlive: 30 * time.Second,
|
|
}
|
|
return newSafeHTTPClient(cfg, net.DefaultResolver.LookupIPAddr, dialer.DialContext)
|
|
}
|
|
|
|
func newSafeHTTPClient(cfg SSRFConfig, lookupIPAddr func(context.Context, string) ([]net.IPAddr, error), dialContext func(context.Context, string, string) (net.Conn, error)) *SafeHTTPClient {
|
|
transport := &http.Transport{
|
|
DialContext: func(ctx context.Context, network, addr string) (net.Conn, error) {
|
|
host, port, err := net.SplitHostPort(addr)
|
|
if err != nil {
|
|
return nil, fmt.Errorf("invalid address: %s", addr)
|
|
}
|
|
|
|
ips, err := lookupIPAddr(ctx, host)
|
|
if err != nil {
|
|
return nil, fmt.Errorf("DNS resolution failed for %s: %w", host, err)
|
|
}
|
|
|
|
for _, ip := range ips {
|
|
if isBlockedIP(ip.IP, cfg.BlockedCIDRs) {
|
|
return nil, fmt.Errorf("SSRF blocked: %s resolves to private IP %s", host, ip.IP)
|
|
}
|
|
}
|
|
|
|
// Dial the validated address directly. Resolving the hostname again here
|
|
// would reopen a DNS-rebinding window between validation and connect.
|
|
var lastErr error
|
|
for _, ip := range ips {
|
|
conn, err := dialContext(ctx, network, net.JoinHostPort(ip.IP.String(), port))
|
|
if err == nil {
|
|
return conn, nil
|
|
}
|
|
lastErr = err
|
|
}
|
|
if lastErr == nil {
|
|
lastErr = fmt.Errorf("DNS resolution returned no addresses for %s", host)
|
|
}
|
|
return nil, lastErr
|
|
},
|
|
MaxIdleConns: 100,
|
|
IdleConnTimeout: 90 * time.Second,
|
|
TLSHandshakeTimeout: 10 * time.Second,
|
|
ResponseHeaderTimeout: 10 * time.Second,
|
|
}
|
|
|
|
return &SafeHTTPClient{
|
|
client: &http.Client{
|
|
Transport: transport,
|
|
Timeout: 30 * time.Second,
|
|
CheckRedirect: safeRedirectCheckConfig(cfg),
|
|
},
|
|
cfg: cfg,
|
|
}
|
|
}
|
|
|
|
// Do executes an HTTP request with SSRF protection.
|
|
func (c *SafeHTTPClient) Do(req *http.Request) (*http.Response, error) {
|
|
if err := validateURLTarget(req.URL, c.cfg); err != nil {
|
|
return nil, err
|
|
}
|
|
|
|
return c.client.Do(req)
|
|
}
|
|
|
|
func isBlockedIP(ip net.IP, blockedCIDRs []string) bool {
|
|
for _, cidr := range blockedCIDRs {
|
|
_, network, err := net.ParseCIDR(cidr)
|
|
if err != nil {
|
|
continue
|
|
}
|
|
if network.Contains(ip) {
|
|
return true
|
|
}
|
|
}
|
|
return false
|
|
}
|
|
|
|
func safeRedirectCheck(maxRedirects int) func(req *http.Request, via []*http.Request) error {
|
|
return func(req *http.Request, via []*http.Request) error {
|
|
if len(via) >= maxRedirects {
|
|
return fmt.Errorf("SSRF: stopped after %d redirects", maxRedirects)
|
|
}
|
|
host := req.URL.Hostname()
|
|
if net.ParseIP(host) != nil {
|
|
return fmt.Errorf("SSRF: redirect to IP literal blocked: %s", host)
|
|
}
|
|
return nil
|
|
}
|
|
}
|
|
|
|
func safeRedirectCheckConfig(cfg SSRFConfig) func(req *http.Request, via []*http.Request) error {
|
|
return func(req *http.Request, via []*http.Request) error {
|
|
if len(via) >= cfg.MaxRedirects {
|
|
return fmt.Errorf("SSRF: stopped after %d redirects", cfg.MaxRedirects)
|
|
}
|
|
return validateURLTarget(req.URL, cfg)
|
|
}
|
|
}
|
|
|
|
// ValidateURL checks if a URL is safe to fetch (without making a request).
|
|
func ValidateURL(rawURL string, cfg SSRFConfig) error {
|
|
if strings.TrimSpace(rawURL) == "" {
|
|
return fmt.Errorf("empty URL")
|
|
}
|
|
parsed, err := url.Parse(rawURL)
|
|
if err != nil {
|
|
return fmt.Errorf("invalid URL: %w", err)
|
|
}
|
|
return validateURLTarget(parsed, cfg)
|
|
}
|
|
|
|
func validateURLTarget(parsed *url.URL, cfg SSRFConfig) error {
|
|
if parsed == nil || parsed.Hostname() == "" {
|
|
return fmt.Errorf("invalid URL: host is required")
|
|
}
|
|
if parsed.Scheme != "http" && parsed.Scheme != "https" {
|
|
return fmt.Errorf("SSRF: unsupported URL scheme %q", parsed.Scheme)
|
|
}
|
|
if cfg.RequireTLS && parsed.Scheme != "https" {
|
|
return fmt.Errorf("SSRF protection: non-HTTPS request blocked for %s", parsed)
|
|
}
|
|
if parsed.User != nil {
|
|
return fmt.Errorf("SSRF: URL credentials are not allowed")
|
|
}
|
|
port := parsed.Port()
|
|
if port == "" {
|
|
if parsed.Scheme == "https" {
|
|
port = "443"
|
|
} else {
|
|
port = "80"
|
|
}
|
|
}
|
|
if len(cfg.AllowedPorts) > 0 {
|
|
allowed := false
|
|
for _, candidate := range cfg.AllowedPorts {
|
|
if port == candidate {
|
|
allowed = true
|
|
break
|
|
}
|
|
}
|
|
if !allowed {
|
|
return fmt.Errorf("SSRF: URL port %s is not allowed", port)
|
|
}
|
|
}
|
|
host := parsed.Hostname()
|
|
if ip := net.ParseIP(host); ip != nil && isBlockedIP(ip, cfg.BlockedCIDRs) {
|
|
return fmt.Errorf("SSRF: URL points to private/reserved IP %s", host)
|
|
}
|
|
return nil
|
|
}
|
|
|
|
// SafeFetchURL fetches a URL safely with SSRF protection.
|
|
func (c *SafeHTTPClient) SafeFetchURL(ctx context.Context, rawurl string) (*http.Response, error) {
|
|
req, err := http.NewRequestWithContext(ctx, "GET", rawurl, nil)
|
|
if err != nil {
|
|
return nil, fmt.Errorf("invalid request: %w", err)
|
|
}
|
|
return c.Do(req)
|
|
}
|