Files
ITWorx-Pulse-Public/internal/probe/policy.go
T
ITWorx Pulse release export bd774932d5
Public source validation / validate (push) Failing after 3m8s
Publish ITWorx Pulse source
2026-09-03 02:09:19 +02:00

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
}