Public source validation / validate (push) Failing after 3m8s
326 lines
9.9 KiB
Go
326 lines
9.9 KiB
Go
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
|
|
}
|