Files
gochat/backend/internal/security/ssrf_protection.go
T
2026-08-22 21:19:43 +08:00

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)
}