diff --git a/pkg/agent/tools/browser/download_pdf.go b/pkg/agent/tools/browser/download_pdf.go index 1a37ec7c3..1a6493c08 100644 --- a/pkg/agent/tools/browser/download_pdf.go +++ b/pkg/agent/tools/browser/download_pdf.go @@ -27,8 +27,8 @@ import ( "github.com/pdfcpu/pdfcpu/pkg/api" "github.com/pdfcpu/pdfcpu/pkg/pdfcpu/model" + "go.gearno.de/kit/httpclient" "go.probo.inc/probo/pkg/agent" - "go.probo.inc/probo/pkg/agent/tools/internal/netcheck" ) type ( @@ -44,10 +44,8 @@ type ( ) func DownloadPDFTool() agent.Tool { - client := &http.Client{ - Timeout: 30 * time.Second, - Transport: netcheck.NewPinnedTransport(), - } + client := httpclient.DefaultPooledClient(httpclient.WithSSRFProtection()) + client.Timeout = 30 * time.Second return agent.FunctionTool( "download_pdf", diff --git a/pkg/agent/tools/browser/fetch_robots.go b/pkg/agent/tools/browser/fetch_robots.go index b18f33667..4d1a1c8d7 100644 --- a/pkg/agent/tools/browser/fetch_robots.go +++ b/pkg/agent/tools/browser/fetch_robots.go @@ -23,6 +23,7 @@ import ( "strings" "time" + "go.gearno.de/kit/httpclient" "go.probo.inc/probo/pkg/agent" ) @@ -40,7 +41,8 @@ type ( ) func FetchRobotsTxtTool() agent.Tool { - client := &http.Client{Timeout: 10 * time.Second} + client := httpclient.DefaultPooledClient(httpclient.WithSSRFProtection()) + client.Timeout = 10 * time.Second return agent.FunctionTool( "fetch_robots_txt", diff --git a/pkg/agent/tools/browser/fetch_sitemap.go b/pkg/agent/tools/browser/fetch_sitemap.go index 8d55c9a22..6414c1b73 100644 --- a/pkg/agent/tools/browser/fetch_sitemap.go +++ b/pkg/agent/tools/browser/fetch_sitemap.go @@ -24,6 +24,7 @@ import ( "strings" "time" + "go.gearno.de/kit/httpclient" "go.probo.inc/probo/pkg/agent" ) @@ -45,7 +46,8 @@ const ( ) func FetchSitemapTool() agent.Tool { - client := &http.Client{Timeout: 15 * time.Second} + client := httpclient.DefaultPooledClient(httpclient.WithSSRFProtection()) + client.Timeout = 15 * time.Second return agent.FunctionTool( "fetch_sitemap", diff --git a/pkg/agent/tools/internal/netcheck/netcheck.go b/pkg/agent/tools/internal/netcheck/netcheck.go index 9b6944253..6e078d1f7 100644 --- a/pkg/agent/tools/internal/netcheck/netcheck.go +++ b/pkg/agent/tools/internal/netcheck/netcheck.go @@ -17,10 +17,8 @@ package netcheck import ( - "context" "fmt" "net" - "net/http" "net/url" ) @@ -88,41 +86,3 @@ func ValidatePublicDomain(domain string) error { return nil } - -// NewPinnedTransport returns an *http.Transport with a custom DialContext that -// resolves the target host once, validates all resolved IPs with IsPublicIP, -// and dials the validated IP directly. This prevents DNS rebinding attacks -// where the first lookup returns a public IP but a subsequent lookup (at -// connection time) returns a private IP. -func NewPinnedTransport() *http.Transport { - return &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("cannot parse address: %w", err) - } - - ips, err := net.DefaultResolver.LookupIPAddr(ctx, host) - if err != nil { - return nil, fmt.Errorf("cannot resolve host: %w", err) - } - - if len(ips) == 0 { - return nil, fmt.Errorf("cannot resolve host: no addresses found") - } - - for _, ip := range ips { - if !IsPublicIP(ip.IP) { - return nil, fmt.Errorf("cannot connect to non-public IP %s", ip.IP) - } - } - - // Dial the first validated IP directly to prevent DNS rebinding. - pinnedAddr := net.JoinHostPort(ips[0].IP.String(), port) - - var d net.Dialer - - return d.DialContext(ctx, network, pinnedAddr) - }, - } -} diff --git a/pkg/agent/tools/security/cors.go b/pkg/agent/tools/security/cors.go index 5f931e5d2..aabc20a6c 100644 --- a/pkg/agent/tools/security/cors.go +++ b/pkg/agent/tools/security/cors.go @@ -21,6 +21,7 @@ import ( "strings" "time" + "go.gearno.de/kit/httpclient" "go.probo.inc/probo/pkg/agent" "go.probo.inc/probo/pkg/agent/tools/internal/netcheck" ) @@ -63,6 +64,12 @@ func splitTrimmed(s, sep string) []string { } func CheckCORSTool() agent.Tool { + client := httpclient.DefaultPooledClient(httpclient.WithSSRFProtection()) + client.Timeout = 10 * time.Second + client.CheckRedirect = func(_ *http.Request, _ []*http.Request) error { + return http.ErrUseLastResponse + } + return agent.FunctionTool( "check_cors", "Send a CORS preflight (OPTIONS) request to a URL with a given Origin and analyze the Access-Control-* response headers, flagging wildcard origins and origin reflection.", @@ -75,13 +82,6 @@ func CheckCORSTool() agent.Tool { ), nil } - client := &http.Client{ - Timeout: 10 * time.Second, - CheckRedirect: func(_ *http.Request, _ []*http.Request) error { - return http.ErrUseLastResponse - }, - } - req, err := http.NewRequestWithContext( ctx, http.MethodOptions, diff --git a/pkg/agent/tools/security/csp.go b/pkg/agent/tools/security/csp.go index b8fc0f94d..16967582d 100644 --- a/pkg/agent/tools/security/csp.go +++ b/pkg/agent/tools/security/csp.go @@ -21,7 +21,9 @@ import ( "strings" "time" + "go.gearno.de/kit/httpclient" "go.probo.inc/probo/pkg/agent" + "go.probo.inc/probo/pkg/agent/tools/internal/netcheck" ) type ( @@ -73,11 +75,20 @@ func parseCSPDirectives(raw string) []cspDirective { } func AnalyzeCSPTool() agent.Tool { + client := httpclient.DefaultPooledClient(httpclient.WithSSRFProtection()) + client.Timeout = 10 * time.Second + return agent.FunctionTool( "analyze_csp", "Analyze the Content-Security-Policy header for a URL, parsing directives and flagging unsafe patterns like unsafe-eval, unsafe-inline, and wildcard sources.", func(ctx context.Context, p cspParams) (agent.ToolResult, error) { - client := &http.Client{Timeout: 10 * time.Second} + if err := netcheck.ValidatePublicURL(p.URL); err != nil { + return agent.ResultJSON( + cspResult{ + ErrorDetail: fmt.Sprintf("URL not allowed: %s", err), + }, + ), nil + } req, err := http.NewRequestWithContext(ctx, http.MethodGet, p.URL, nil) if err != nil { diff --git a/pkg/agent/tools/security/headers.go b/pkg/agent/tools/security/headers.go index a358ec72f..aae457b2d 100644 --- a/pkg/agent/tools/security/headers.go +++ b/pkg/agent/tools/security/headers.go @@ -22,6 +22,7 @@ import ( "strings" "time" + "go.gearno.de/kit/httpclient" "go.probo.inc/probo/pkg/agent" "go.probo.inc/probo/pkg/agent/tools/internal/netcheck" ) @@ -75,6 +76,15 @@ func headersFromResponse(resp *http.Response) headersResult { } func CheckSecurityHeadersTool() agent.Tool { + client := httpclient.DefaultPooledClient(httpclient.WithSSRFProtection()) + client.Timeout = 10 * time.Second + client.CheckRedirect = func(_ *http.Request, _ []*http.Request) error { + return http.ErrUseLastResponse + } + + followClient := httpclient.DefaultPooledClient(httpclient.WithSSRFProtection()) + followClient.Timeout = 10 * time.Second + return agent.FunctionTool( "check_security_headers", "Check security-related HTTP headers for a URL (HSTS, CSP, X-Frame-Options, X-Content-Type-Options, Referrer-Policy, Permissions-Policy, Cross-Origin-*-Policy). Also checks if HTTP redirects to HTTPS.", @@ -87,13 +97,6 @@ func CheckSecurityHeadersTool() agent.Tool { ), nil } - client := &http.Client{ - Timeout: 10 * time.Second, - CheckRedirect: func(_ *http.Request, _ []*http.Request) error { - return http.ErrUseLastResponse - }, - } - // First check the HTTP version to detect HTTP→HTTPS redirect. redirectsToHTTPS := false @@ -129,8 +132,6 @@ func CheckSecurityHeadersTool() agent.Tool { httpsParsed.Scheme = "https" httpsURL := httpsParsed.String() - followClient := &http.Client{Timeout: 10 * time.Second} - httpsReq, err := http.NewRequestWithContext(ctx, http.MethodGet, httpsURL, nil) if err != nil { return agent.ResultJSON(