package probe import ( "context" "errors" "fmt" "io" "net" "net/http" "net/netip" "net/url" "strconv" "strings" "time" "github.com/itworx/pulse/internal/audit" ) const ( DefaultMaxResponseBytes int64 = 1 << 20 DefaultMaxHeaderBytes = 32 << 10 DefaultMaxHeaders = 64 DefaultMaxRedirects = 5 ) type Resolver interface { LookupIP(context.Context, string) ([]netip.Addr, error) } type NetResolver struct{} func (NetResolver) LookupIP(ctx context.Context, host string) ([]netip.Addr, error) { return net.DefaultResolver.LookupNetIP(ctx, "ip", host) } type NetworkPolicy struct { ID string Revision int64 AllowedNetworks []netip.Prefix AllowedSchemes []string AllowedMethods []string AllowedHeaders []string MaxResponseBytes int64 MaxHeaderBytes int MaxHeaders int MaxRedirects int Timeout time.Duration } func (p NetworkPolicy) withDefaults() NetworkPolicy { if p.Revision == 0 { p.Revision = 1 } if len(p.AllowedSchemes) == 0 { p.AllowedSchemes = []string{"http", "https", "tcp", "dns", "icmp", "tls"} } if len(p.AllowedMethods) == 0 { p.AllowedMethods = []string{http.MethodGet, http.MethodHead} } if len(p.AllowedHeaders) == 0 { p.AllowedHeaders = []string{"Accept", "User-Agent"} } if p.MaxResponseBytes == 0 { p.MaxResponseBytes = DefaultMaxResponseBytes } if p.MaxHeaderBytes == 0 { p.MaxHeaderBytes = DefaultMaxHeaderBytes } if p.MaxHeaders == 0 { p.MaxHeaders = DefaultMaxHeaders } if p.MaxRedirects == 0 { p.MaxRedirects = DefaultMaxRedirects } if p.Timeout == 0 { p.Timeout = 10 * time.Second } return p } func (p NetworkPolicy) Validate() error { p = p.withDefaults() if p.Revision < 1 || len(p.AllowedNetworks) > 32 || p.MaxResponseBytes < 1 || p.MaxResponseBytes > 16<<20 || p.MaxHeaderBytes < 1024 || p.MaxHeaderBytes > 256<<10 || p.MaxHeaders < 1 || p.MaxHeaders > 256 || p.MaxRedirects < 0 || p.MaxRedirects > 10 || p.Timeout <= 0 || p.Timeout > 2*time.Minute { return errors.New("network policy is outside safe bounds") } for _, prefix := range p.AllowedNetworks { if !prefix.IsValid() { return errors.New("network policy contains an invalid allowlist prefix") } } if len(p.AllowedSchemes) == 0 || len(p.AllowedMethods) == 0 { return errors.New("network policy requires schemes and methods") } return nil } func (p NetworkPolicy) ValidateTarget(ctx context.Context, resolver Resolver, scheme, host string, port int) ([]netip.Addr, error) { p = p.withDefaults() if err := p.Validate(); err != nil { return nil, err } if ctx == nil { return nil, errors.New("target context is nil") } if !containsFold(p.AllowedSchemes, scheme) { return nil, errors.New("probe scheme is not allowed") } return p.validateAddress(ctx, resolver, host, port) } // validateAddress repeats the address safety checks at dial time without // coupling the transport callback to a particular application scheme. func (p NetworkPolicy) validateAddress(ctx context.Context, resolver Resolver, host string, port int) ([]netip.Addr, error) { if ctx == nil { return nil, errors.New("target context is nil") } if strings.TrimSpace(host) == "" || strings.ContainsAny(host, "/?#@") || len(host) > 253 || port < 1 || port > 65535 { return nil, errors.New("probe target host or port is invalid") } if resolver == nil { resolver = NetResolver{} } if ip, err := netip.ParseAddr(host); err == nil { if !p.addressAllowed(ip) { return nil, fmt.Errorf("probe target address %s is blocked", ip) } return []netip.Addr{ip}, nil } addresses, err := resolver.LookupIP(ctx, strings.TrimSuffix(strings.ToLower(host), ".")) if err != nil { return nil, fmt.Errorf("resolve probe target: %w", err) } if len(addresses) == 0 { return nil, errors.New("probe target has no addresses") } for _, address := range addresses { if !p.addressAllowed(address) { return nil, fmt.Errorf("resolved probe target address %s is blocked", address) } } return addresses, nil } func (p NetworkPolicy) addressAllowed(address netip.Addr) bool { if !address.IsValid() || address.IsUnspecified() || address.IsLinkLocalUnicast() || address.IsMulticast() || isMetadataAddress(address) { return false } if !address.IsPrivate() && !address.IsLoopback() { return true } for _, prefix := range p.AllowedNetworks { if prefix.Contains(address) { return true } } return false } func isMetadataAddress(address netip.Addr) bool { return address == netip.MustParseAddr("169.254.169.254") || address == netip.MustParseAddr("fd00:ec2::254") } func (p NetworkPolicy) ValidateURL(ctx context.Context, resolver Resolver, raw string) (*url.URL, error) { parsed, err := url.Parse(raw) if err != nil || parsed.Scheme == "" || parsed.Hostname() == "" || parsed.User != nil || parsed.Fragment != "" { return nil, errors.New("probe URL is invalid") } port := 0 if rawPort := parsed.Port(); rawPort != "" { port, err = strconv.Atoi(rawPort) if err != nil { return nil, errors.New("probe URL port is invalid") } } else if parsed.Scheme == "http" { port = 80 } else if parsed.Scheme == "https" { port = 443 } if port == 0 { return nil, errors.New("probe URL requires a port") } if _, err := p.ValidateTarget(ctx, resolver, parsed.Scheme, parsed.Hostname(), port); err != nil { return nil, err } return parsed, nil } func (p NetworkPolicy) ValidateRequest(request *http.Request) error { p = p.withDefaults() if request == nil || request.URL == nil || !containsFold(p.AllowedMethods, request.Method) || request.Body != nil || request.ContentLength > 0 { return errors.New("probe request method or body is not allowed") } if len(request.Header) > p.MaxHeaders { return errors.New("probe request has too many headers") } for name, values := range request.Header { if !containsFold(p.AllowedHeaders, name) { return fmt.Errorf("probe request header %q is not allowed", name) } for _, value := range values { if len(value) > p.MaxHeaderBytes { return errors.New("probe request header is too large") } } } return nil } type SafeClient struct { HTTP *http.Client Policy NetworkPolicy Resolver Resolver } func NewSafeClient(policy NetworkPolicy, resolver Resolver) (*SafeClient, error) { policy = policy.withDefaults() if err := policy.Validate(); err != nil { return nil, err } if resolver == nil { resolver = NetResolver{} } dialer := &net.Dialer{Timeout: policy.Timeout} transport := &http.Transport{Proxy: nil, MaxResponseHeaderBytes: int64(policy.MaxHeaderBytes)} transport.DialContext = func(ctx context.Context, network, address string) (net.Conn, error) { host, port, err := net.SplitHostPort(address) if err != nil { return nil, err } addresses, err := policy.validateAddress(ctx, resolver, host, mustPort(port)) if err != nil { return nil, err } var lastErr error for _, resolved := range addresses { connection, dialErr := dialer.DialContext(ctx, network, net.JoinHostPort(resolved.String(), port)) if dialErr == nil { return connection, nil } lastErr = dialErr } return nil, lastErr } client := &http.Client{Transport: transport, Timeout: policy.Timeout} client.CheckRedirect = func(request *http.Request, via []*http.Request) error { // Never carry the prior URL as a referrer across probe targets. request.Header.Del("Referer") if len(via) > policy.MaxRedirects { return errors.New("probe redirect limit exceeded") } if err := policy.ValidateRequest(request); err != nil { return err } _, err := policy.ValidateURL(request.Context(), resolver, request.URL.String()) return err } return &SafeClient{HTTP: client, Policy: policy, Resolver: resolver}, nil } func (c *SafeClient) Do(request *http.Request) (*http.Response, error) { if c == nil || c.HTTP == nil { return nil, errors.New("safe probe client is nil") } if err := c.Policy.ValidateRequest(request); err != nil { return nil, err } if _, err := c.Policy.ValidateURL(request.Context(), c.Resolver, request.URL.String()); err != nil { return nil, err } response, err := c.HTTP.Do(request) if err != nil { return nil, err } if response.ContentLength > c.Policy.MaxResponseBytes { response.Body.Close() return nil, errors.New("probe response exceeds size limit") } response.Body = io.NopCloser(io.LimitReader(response.Body, c.Policy.MaxResponseBytes+1)) return response, nil } func ReadLimitedBody(response *http.Response, maxBytes int64) ([]byte, error) { if response == nil || response.Body == nil || maxBytes < 1 { return nil, errors.New("probe response is invalid") } defer response.Body.Close() body, err := io.ReadAll(io.LimitReader(response.Body, maxBytes+1)) if err != nil { return nil, err } if int64(len(body)) > maxBytes { return nil, errors.New("probe response exceeds size limit") } return body, nil } func RecordPolicyChange(ctx context.Context, store audit.Store, actor, policyID, correlationID string, before, after NetworkPolicy) error { if store == nil { return errors.New("probe policy audit store is nil") } return store.Append(ctx, audit.Event{Actor: actor, Action: "probe.network_policy.update", ResourceType: "network_policy", ResourceID: policyID, Result: "success", CorrelationID: correlationID, Before: policyAuditDocument(before), After: policyAuditDocument(after)}) } func policyAuditDocument(policy NetworkPolicy) map[string]any { policy = policy.withDefaults() prefixes := make([]string, 0, len(policy.AllowedNetworks)) for _, prefix := range policy.AllowedNetworks { prefixes = append(prefixes, prefix.String()) } return map[string]any{"revision": policy.Revision, "allowedNetworks": prefixes, "allowedSchemes": policy.AllowedSchemes, "allowedMethods": policy.AllowedMethods, "maxResponseBytes": policy.MaxResponseBytes, "maxRedirects": policy.MaxRedirects} } func containsFold(values []string, target string) bool { for _, value := range values { if strings.EqualFold(value, target) { return true } } return false } func mustPort(raw string) int { port, _ := strconv.Atoi(raw) return port }