Publish ITWorx Pulse source
Public source validation / validate (push) Failing after 3m8s

This commit is contained in:
ITWorx Pulse release export
2026-09-03 02:09:19 +02:00
commit bd774932d5
614 changed files with 77116 additions and 0 deletions
+266
View File
@@ -0,0 +1,266 @@
package probe
import (
"context"
"crypto/tls"
"crypto/x509"
"encoding/json"
"errors"
"net"
"net/http"
"net/netip"
"net/url"
"reflect"
"strconv"
"strings"
"time"
)
type HTTPDoer interface {
Do(*http.Request) (*http.Response, error)
}
type ProbeExecutor struct {
Policy NetworkPolicy
Resolver Resolver
HTTP HTTPDoer
DialContext func(context.Context, string, string) (net.Conn, error)
Now func() time.Time
}
func (e ProbeExecutor) Execute(ctx context.Context, definition Definition) (Result, error) {
if ctx == nil {
return Result{ProbeID: definition.ID, State: "unknown", ErrorClass: "invalid_context"}, errors.New("probe context is nil")
}
if err := ctx.Err(); err != nil {
return Result{ProbeID: definition.ID, State: "unknown", ErrorClass: "canceled"}, err
}
if err := definition.Validate(); err != nil {
return Result{ProbeID: definition.ID, State: "unknown", ErrorClass: "invalid_probe"}, err
}
if e.Resolver == nil {
e.Resolver = NetResolver{}
}
if e.Now == nil {
e.Now = func() time.Time { return time.Now().UTC() }
}
probeContext, cancel := context.WithTimeout(ctx, definition.Timeout)
defer cancel()
switch definition.Type {
case TypeHTTP:
return e.executeHTTP(probeContext, definition)
case TypeTCP:
return e.executeTCP(probeContext, definition)
case TypeDNS:
return e.executeDNS(probeContext, definition)
case TypeTLS:
return e.executeTLS(probeContext, definition)
case TypeICMP:
return Result{ProbeID: definition.ID, ObservedAt: e.Now().UTC(), CompletedAt: e.Now().UTC(), State: "unknown", ErrorClass: "unsupported"}, nil
default:
return Result{ProbeID: definition.ID, ObservedAt: e.Now().UTC(), CompletedAt: e.Now().UTC(), State: "unknown", ErrorClass: "unsupported"}, nil
}
}
func (e ProbeExecutor) executeHTTP(ctx context.Context, definition Definition) (Result, error) {
started := e.Now().UTC()
scheme := strings.ToLower(strings.TrimSpace(definition.Target.Scheme))
if scheme == "" {
scheme = "http"
}
port := defaultPort(scheme, 80)
if definition.Target.Port > 0 {
port = definition.Target.Port
}
if _, err := e.Policy.ValidateTarget(ctx, e.Resolver, scheme, definition.Target.Host, port); err != nil {
return transportFailure(definition.ID, err), err
}
target := &url.URL{Scheme: scheme, Host: net.JoinHostPort(definition.Target.Host, strconv.Itoa(port)), Path: definition.Target.Path}
if target.Path == "" {
target.Path = "/"
}
request, err := http.NewRequestWithContext(ctx, http.MethodGet, target.String(), nil)
if err != nil {
return transportFailure(definition.ID, err), err
}
if e.HTTP == nil {
client, clientErr := NewSafeClient(e.Policy, e.Resolver)
if clientErr != nil {
return transportFailure(definition.ID, clientErr), clientErr
}
if !definition.FollowRedirects {
client.HTTP.CheckRedirect = func(*http.Request, []*http.Request) error { return http.ErrUseLastResponse }
}
e.HTTP = client
}
response, err := e.HTTP.Do(request)
if err != nil {
return transportFailure(definition.ID, err), err
}
if response == nil {
err = errors.New("probe response is invalid")
return transportFailure(definition.ID, err), err
}
body, bodyErr := ReadLimitedBody(response, e.Policy.withDefaults().MaxResponseBytes)
result := Result{ProbeID: definition.ID, ObservedAt: started, CompletedAt: e.Now().UTC(), ResponseTimeMS: intPointer(int(time.Since(started).Milliseconds())), StatusCode: intPointer(response.StatusCode)}
if bodyErr != nil {
result.State = "unknown"
result.ErrorClass = "response_too_large"
return result, nil
}
if !expectedStatus(definition.ExpectedStatusCodes, response.StatusCode) {
result.State = "down"
result.ErrorClass = "status_not_expected"
return result, nil
}
if assertionErr := assertBody(body, definition.ContentAssertion); assertionErr != nil {
result.State = "degraded"
result.ErrorClass = "content_assertion_failed"
return result, nil
}
result.State = "up"
return result, nil
}
func (e ProbeExecutor) executeTCP(ctx context.Context, definition Definition) (Result, error) {
started := e.Now().UTC()
port := definition.Target.Port
if port == 0 {
return transportFailure(definition.ID, errors.New("tcp port is required")), errors.New("tcp port is required")
}
addresses, err := e.Policy.ValidateTarget(ctx, e.Resolver, TypeTCP, definition.Target.Host, port)
if err != nil {
return transportFailure(definition.ID, err), err
}
connection, err := e.dial(ctx, "tcp", addresses[0], port)
if err != nil {
return transportFailure(definition.ID, err), err
}
if connection == nil {
err = errors.New("probe dial returned nil connection")
return transportFailure(definition.ID, err), err
}
_ = connection.Close()
return Result{ProbeID: definition.ID, ObservedAt: started, CompletedAt: e.Now().UTC(), State: "up", ResponseTimeMS: intPointer(int(time.Since(started).Milliseconds()))}, nil
}
func (e ProbeExecutor) executeDNS(ctx context.Context, definition Definition) (Result, error) {
started := e.Now().UTC()
addresses, err := e.Policy.ValidateTarget(ctx, e.Resolver, TypeDNS, definition.Target.Host, defaultPort("dns", 53))
if err != nil {
return transportFailure(definition.ID, err), err
}
return Result{ProbeID: definition.ID, ObservedAt: started, CompletedAt: e.Now().UTC(), State: "up", ResponseTimeMS: intPointer(int(time.Since(started).Milliseconds())), Attributes: map[string]any{"addressCount": len(addresses)}}, nil
}
func (e ProbeExecutor) executeTLS(ctx context.Context, definition Definition) (Result, error) {
started := e.Now().UTC()
port := defaultPort(definition.Target.Scheme, 443)
if definition.Target.Port > 0 {
port = definition.Target.Port
}
addresses, err := e.Policy.ValidateTarget(ctx, e.Resolver, TypeTLS, definition.Target.Host, port)
if err != nil {
return transportFailure(definition.ID, err), err
}
connection, err := e.dial(ctx, "tcp", addresses[0], port)
if err != nil {
return transportFailure(definition.ID, err), err
}
if connection == nil {
err = errors.New("probe dial returned nil connection")
return transportFailure(definition.ID, err), err
}
defer connection.Close()
tlsConnection := tls.Client(connection, &tls.Config{ServerName: definition.Target.Host, MinVersion: tls.VersionTLS12, InsecureSkipVerify: !definition.VerifyTLS})
if err := tlsConnection.HandshakeContext(ctx); err != nil {
return transportFailure(definition.ID, err), err
}
state := tlsConnection.ConnectionState()
result := Result{ProbeID: definition.ID, ObservedAt: started, CompletedAt: e.Now().UTC(), State: "up", ResponseTimeMS: intPointer(int(time.Since(started).Milliseconds()))}
if len(state.PeerCertificates) > 0 {
result.Certificate = certificateFromX509(state.PeerCertificates[0], definition, e.Now)
if result.Certificate.VerificationState == "attention" || result.Certificate.VerificationState == "invalid" {
result.State = "degraded"
}
}
return result, nil
}
func (e ProbeExecutor) dial(ctx context.Context, network string, address netip.Addr, port int) (net.Conn, error) {
if e.DialContext != nil {
return e.DialContext(ctx, network, net.JoinHostPort(address.String(), strconv.Itoa(port)))
}
dialer := &net.Dialer{Timeout: e.Policy.withDefaults().Timeout}
return dialer.DialContext(ctx, network, net.JoinHostPort(address.String(), strconv.Itoa(port)))
}
func certificateFromX509(certificate *x509.Certificate, definition Definition, now func() time.Time) *Certificate {
if certificate == nil {
return nil
}
observed := now().UTC()
hostnameValid := certificate.VerifyHostname(definition.Target.Host) == nil
verification := "valid"
if !hostnameValid {
verification = "invalid"
} else if !certificate.NotAfter.After(observed) {
verification = "invalid"
} else if certificate.NotAfter.Before(observed.Add(30 * 24 * time.Hour)) {
verification = "attention"
}
return &Certificate{ID: definition.ID + "-certificate", ServiceID: definition.ServiceID, EndpointID: definition.EndpointID, ObservedAt: observed, ExpiresAt: &certificate.NotAfter, Issuer: certificate.Issuer.String(), Subject: certificate.Subject.String(), HostnameValid: &hostnameValid, VerificationState: verification}
}
func assertBody(body []byte, assertion map[string]any) error {
if len(assertion) == 0 {
return nil
}
if keyword, ok := assertion["keyword"].(string); ok && !strings.Contains(string(body), keyword) {
return errors.New("keyword assertion failed")
}
if expected, ok := assertion["json"].(map[string]any); ok {
var actual map[string]any
if err := json.Unmarshal(body, &actual); err != nil {
return errors.New("json assertion body is invalid")
}
for key, value := range expected {
if !reflect.DeepEqual(actual[key], value) {
return errors.New("json assertion failed")
}
}
}
return nil
}
func expectedStatus(expected []int, status int) bool {
if len(expected) == 0 {
return status >= 200 && status < 400
}
for _, candidate := range expected {
if candidate == status {
return true
}
}
return false
}
func defaultPort(scheme string, fallback int) int {
if scheme == "http" {
return 80
}
if scheme == "https" || scheme == "tls" {
return 443
}
return fallback
}
func transportFailure(id string, err error) Result {
result := Result{ProbeID: id, State: "unknown", ErrorClass: "transport_error"}
if err != nil {
result.ErrorMessage = boundText(err.Error(), 256)
}
return result
}
func intPointer(value int) *int { return &value }
+246
View File
@@ -0,0 +1,246 @@
package probe
import (
"context"
"crypto/rand"
"crypto/rsa"
"crypto/tls"
"crypto/x509"
"crypto/x509/pkix"
"errors"
"io"
"math/big"
"net"
"net/http"
"net/http/httptest"
"net/netip"
"net/url"
"strings"
"testing"
"time"
)
type fakeHTTPDoer struct {
response *http.Response
err error
request *http.Request
}
func (f *fakeHTTPDoer) Do(request *http.Request) (*http.Response, error) {
f.request = request
return f.response, f.err
}
func probeDefinition(probeType string) Definition {
return Definition{
ID: "probe-1",
ServiceID: "service-1",
Name: "probe",
Type: probeType,
Target: Target{Scheme: "http", Host: "service.internal", Port: 8080, Path: "/health"},
Interval: 30 * time.Second,
Timeout: 2 * time.Second,
Enabled: true,
Revision: 1,
VerifyTLS: false,
FollowRedirects: false,
}
}
func executorPolicy() NetworkPolicy {
return NetworkPolicy{Revision: 1, AllowedNetworks: []netip.Prefix{netip.MustParsePrefix("10.0.0.0/8"), netip.MustParsePrefix("127.0.0.0/8")}}
}
func executorResolver(_ context.Context, _ string) ([]netip.Addr, error) {
return []netip.Addr{netip.MustParseAddr("10.10.0.8")}, nil
}
func TestProbeExecutorHTTPAssertionsAndStatus(t *testing.T) {
body := `{"ok":true,"message":"ready"}`
doer := &fakeHTTPDoer{response: &http.Response{StatusCode: http.StatusOK, Body: io.NopCloser(strings.NewReader(body))}}
definition := probeDefinition(TypeHTTP)
definition.ContentAssertion = map[string]any{"keyword": "ready", "json": map[string]any{"ok": true}}
executor := ProbeExecutor{Policy: executorPolicy(), Resolver: resolverFunc(executorResolver), HTTP: doer}
result, err := executor.Execute(context.Background(), definition)
if err != nil {
t.Fatal(err)
}
if result.State != "up" || result.StatusCode == nil || *result.StatusCode != http.StatusOK {
t.Fatalf("unexpected HTTP result: %+v", result)
}
if doer.request == nil || doer.request.Method != http.MethodGet || doer.request.URL.String() != "http://service.internal:8080/health" {
t.Fatalf("unexpected request: %+v", doer.request)
}
}
func TestProbeExecutorHTTPFailureBodyLimitAndRedaction(t *testing.T) {
definition := probeDefinition(TypeHTTP)
definition.ExpectedStatusCodes = []int{http.StatusOK}
definition.SecretReference = "secret-ref-should-never-leak"
failureDoer := &fakeHTTPDoer{response: &http.Response{StatusCode: http.StatusServiceUnavailable, Body: io.NopCloser(strings.NewReader("down"))}}
failure, err := (ProbeExecutor{Policy: executorPolicy(), Resolver: resolverFunc(executorResolver), HTTP: failureDoer}).Execute(context.Background(), definition)
if err != nil {
t.Fatal(err)
}
if failure.State != "down" || failure.ErrorClass != "status_not_expected" || strings.Contains(failure.ErrorMessage, definition.SecretReference) {
t.Fatalf("unexpected failure result: %+v", failure)
}
if failureDoer.request == nil || failureDoer.request.Header.Get("Authorization") != "" {
t.Fatal("probe credentials were placed in the request")
}
limitedPolicy := executorPolicy()
limitedPolicy.MaxResponseBytes = 4
largeDoer := &fakeHTTPDoer{response: &http.Response{StatusCode: http.StatusOK, Body: io.NopCloser(strings.NewReader("12345"))}}
large, err := (ProbeExecutor{Policy: limitedPolicy, Resolver: resolverFunc(executorResolver), HTTP: largeDoer}).Execute(context.Background(), definition)
if err != nil {
t.Fatal(err)
}
if large.State != "unknown" || large.ErrorClass != "response_too_large" {
t.Fatalf("unexpected bounded response result: %+v", large)
}
}
func TestProbeExecutorHTTPRedirectPolicyAndTimeout(t *testing.T) {
server := httptest.NewServer(http.HandlerFunc(func(response http.ResponseWriter, request *http.Request) {
if request.URL.Path == "/start" {
http.Redirect(response, request, "/final", http.StatusFound)
return
}
_, _ = response.Write([]byte("final"))
}))
defer server.Close()
serverURL, err := url.Parse(server.URL)
if err != nil {
t.Fatal(err)
}
port := mustPort(serverURL.Port())
resolver := resolverFunc(func(_ context.Context, host string) ([]netip.Addr, error) {
if host == "service.internal" {
return []netip.Addr{netip.MustParseAddr("127.0.0.1")}, nil
}
return nil, errors.New("unexpected host")
})
definition := probeDefinition(TypeHTTP)
definition.Target = Target{Scheme: "http", Host: "service.internal", Port: port, Path: "/start"}
definition.FollowRedirects = false
withoutRedirect, err := (ProbeExecutor{Policy: executorPolicy(), Resolver: resolver}).Execute(context.Background(), definition)
if err != nil {
t.Fatal(err)
}
if withoutRedirect.State != "up" || withoutRedirect.StatusCode == nil || *withoutRedirect.StatusCode != http.StatusFound {
t.Fatalf("redirect was not safely stopped: %+v", withoutRedirect)
}
definition.FollowRedirects = true
withRedirect, err := (ProbeExecutor{Policy: executorPolicy(), Resolver: resolver}).Execute(context.Background(), definition)
if err != nil {
t.Fatal(err)
}
if withRedirect.State != "up" || withRedirect.StatusCode == nil || *withRedirect.StatusCode != http.StatusOK {
t.Fatalf("redirect was not followed within policy: %+v", withRedirect)
}
}
func TestProbeExecutorTimeoutAndDNS(t *testing.T) {
definition := probeDefinition(TypeHTTP)
definition.Timeout = time.Second
// Use a doer that releases only when the executor's per-probe context expires.
blockingDoer := httpDoerFunc(func(request *http.Request) (*http.Response, error) {
<-request.Context().Done()
return nil, request.Context().Err()
})
result, err := (ProbeExecutor{Policy: executorPolicy(), Resolver: resolverFunc(executorResolver), HTTP: blockingDoer}).Execute(context.Background(), definition)
if !errors.Is(err, context.DeadlineExceeded) || result.ErrorClass != "transport_error" {
t.Fatalf("unexpected timeout result=%+v err=%v", result, err)
}
dnsDefinition := probeDefinition(TypeDNS)
dnsDefinition.Target = Target{Host: "dns.internal"}
dnsResolver := resolverFunc(func(_ context.Context, host string) ([]netip.Addr, error) {
if host != "dns.internal" {
return nil, errors.New("unexpected host")
}
return []netip.Addr{netip.MustParseAddr("10.0.0.1"), netip.MustParseAddr("10.0.0.2")}, nil
})
dnsResult, err := (ProbeExecutor{Policy: executorPolicy(), Resolver: dnsResolver}).Execute(context.Background(), dnsDefinition)
if err != nil || dnsResult.State != "up" || dnsResult.Attributes["addressCount"] != 2 {
t.Fatalf("unexpected DNS result=%+v err=%v", dnsResult, err)
}
}
func TestProbeExecutorTCPAndUnsupportedICMP(t *testing.T) {
definition := probeDefinition(TypeTCP)
definition.Target = Target{Host: "tcp.internal", Port: 443}
calledNetwork := ""
tcpResult, err := (ProbeExecutor{Policy: executorPolicy(), Resolver: resolverFunc(executorResolver), DialContext: func(_ context.Context, network, _ string) (net.Conn, error) {
calledNetwork = network
client, peer := net.Pipe()
_ = peer.Close()
return client, nil
}}).Execute(context.Background(), definition)
if err != nil || tcpResult.State != "up" || calledNetwork != "tcp" {
t.Fatalf("unexpected TCP result=%+v err=%v network=%s", tcpResult, err, calledNetwork)
}
icmpDefinition := probeDefinition(TypeICMP)
icmpDefinition.Target = Target{Host: "gateway.internal"}
icmpResult, err := (ProbeExecutor{Policy: executorPolicy()}).Execute(context.Background(), icmpDefinition)
if err != nil || icmpResult.State != "unknown" || icmpResult.ErrorClass != "unsupported" {
t.Fatalf("unexpected ICMP fallback result=%+v err=%v", icmpResult, err)
}
}
func TestProbeExecutorTLSCertificateFacts(t *testing.T) {
observed := time.Date(2026, time.August, 2, 12, 0, 0, 0, time.UTC)
privateKey, err := rsa.GenerateKey(rand.Reader, 2048)
if err != nil {
t.Fatal(err)
}
certificateTemplate := &x509.Certificate{SerialNumber: big.NewInt(1), Subject: pkix.Name{CommonName: "tls.internal"}, DNSNames: []string{"tls.internal"}, NotBefore: observed.Add(-time.Hour), NotAfter: observed.Add(90 * 24 * time.Hour), KeyUsage: x509.KeyUsageDigitalSignature, ExtKeyUsage: []x509.ExtKeyUsage{x509.ExtKeyUsageServerAuth}}
certificateDER, err := x509.CreateCertificate(rand.Reader, certificateTemplate, certificateTemplate, &privateKey.PublicKey, privateKey)
if err != nil {
t.Fatal(err)
}
serverCertificate := tls.Certificate{Certificate: [][]byte{certificateDER}, PrivateKey: privateKey}
definition := probeDefinition(TypeTLS)
definition.Target = Target{Scheme: "https", Host: "tls.internal", Port: 443}
definition.VerifyTLS = false
tlsResult, err := (ProbeExecutor{
Policy: executorPolicy(),
Resolver: resolverFunc(func(context.Context, string) ([]netip.Addr, error) {
return []netip.Addr{netip.MustParseAddr("10.0.0.3")}, nil
}),
Now: func() time.Time { return observed },
DialContext: func(_ context.Context, network, _ string) (net.Conn, error) {
if network != "tcp" {
return nil, errors.New("unexpected TLS network")
}
client, server := net.Pipe()
go func() {
defer server.Close()
_ = tls.Server(server, &tls.Config{Certificates: []tls.Certificate{serverCertificate}}).Handshake()
}()
return client, nil
},
}).Execute(context.Background(), definition)
if err != nil || tlsResult.State != "up" || tlsResult.Certificate == nil || tlsResult.Certificate.VerificationState != "valid" || tlsResult.Certificate.HostnameValid == nil || !*tlsResult.Certificate.HostnameValid {
t.Fatalf("unexpected TLS result=%+v err=%v", tlsResult, err)
}
attention := *certificateTemplate
attention.NotAfter = observed.Add(7 * 24 * time.Hour)
attentionResult := certificateFromX509(&attention, definition, func() time.Time { return observed })
if attentionResult.VerificationState != "attention" {
t.Fatalf("expected certificate attention, got %+v", attentionResult)
}
invalid := *certificateTemplate
invalid.DNSNames = []string{"other.internal"}
invalidResult := certificateFromX509(&invalid, definition, func() time.Time { return observed })
if invalidResult.VerificationState != "invalid" || invalidResult.HostnameValid == nil || *invalidResult.HostnameValid {
t.Fatalf("expected invalid hostname, got %+v", invalidResult)
}
}
type httpDoerFunc func(*http.Request) (*http.Response, error)
func (f httpDoerFunc) Do(request *http.Request) (*http.Response, error) { return f(request) }
+325
View File
@@ -0,0 +1,325 @@
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
}
+109
View File
@@ -0,0 +1,109 @@
package probe
import (
"context"
"net/http"
"net/http/httptest"
"net/netip"
"strings"
"testing"
"github.com/itworx/pulse/internal/audit"
)
type resolverFunc func(context.Context, string) ([]netip.Addr, error)
func (f resolverFunc) LookupIP(ctx context.Context, host string) ([]netip.Addr, error) {
return f(ctx, host)
}
func privatePolicy() NetworkPolicy {
return NetworkPolicy{Revision: 1, AllowedNetworks: []netip.Prefix{netip.MustParsePrefix("10.0.0.0/8")}}
}
func TestNetworkPolicyBlocksSSRFAndAllowsExplicitLAN(t *testing.T) {
policy := privatePolicy()
resolver := resolverFunc(func(_ context.Context, host string) ([]netip.Addr, error) {
if host == "lan.internal" {
return []netip.Addr{netip.MustParseAddr("10.10.0.8")}, nil
}
return []netip.Addr{netip.MustParseAddr("169.254.169.254")}, nil
})
if _, err := policy.ValidateTarget(context.Background(), resolver, "https", "lan.internal", 443); err != nil {
t.Fatal(err)
}
if _, err := policy.ValidateTarget(context.Background(), resolver, "https", "metadata.internal", 443); err == nil {
t.Fatal("metadata target was allowed")
}
if _, err := policy.ValidateTarget(context.Background(), resolver, "https", resolverHost("127.0.0.1"), 80); err == nil {
t.Fatal("loopback target was allowed")
}
}
func TestNetworkPolicyRejectsDNSRebindingAndRedirectEscape(t *testing.T) {
policy := privatePolicy()
addresses := []netip.Addr{netip.MustParseAddr("10.0.0.2")}
resolver := resolverFunc(func(_ context.Context, _ string) ([]netip.Addr, error) {
current := addresses
addresses = []netip.Addr{netip.MustParseAddr("169.254.169.254")}
return current, nil
})
if _, err := policy.ValidateTarget(context.Background(), resolver, "https", "service.internal", 443); err != nil {
t.Fatal(err)
}
if _, err := policy.ValidateTarget(context.Background(), resolver, "https", "service.internal", 443); err == nil {
t.Fatal("DNS rebinding to a blocked address was allowed")
}
resolver = resolverFunc(func(_ context.Context, _ string) ([]netip.Addr, error) {
return []netip.Addr{netip.MustParseAddr("10.0.0.2"), netip.MustParseAddr("169.254.169.254")}, nil
})
if _, err := policy.ValidateTarget(context.Background(), resolver, "https", "service.internal", 443); err == nil {
t.Fatal("mixed DNS answer was allowed")
}
client, err := NewSafeClient(policy, resolverFunc(func(_ context.Context, host string) ([]netip.Addr, error) {
if host == "safe.internal" {
return []netip.Addr{netip.MustParseAddr("10.0.0.2")}, nil
}
return []netip.Addr{netip.MustParseAddr("169.254.169.254")}, nil
}))
if err != nil {
t.Fatal(err)
}
redirect := httptest.NewRequest(http.MethodGet, "https://metadata.internal/", nil)
if err := client.HTTP.CheckRedirect(redirect, nil); err == nil {
t.Fatal("redirect escape was allowed")
}
}
func TestNetworkPolicyBoundsRequestsResponsesAndAudits(t *testing.T) {
policy := privatePolicy()
request := httptest.NewRequest(http.MethodPost, "https://example.com", strings.NewReader("body"))
if err := policy.ValidateRequest(request); err == nil {
t.Fatal("POST/body request was allowed")
}
request = httptest.NewRequest(http.MethodGet, "https://example.com", nil)
request.Header.Set("Authorization", "secret")
if err := policy.ValidateRequest(request); err == nil {
t.Fatal("authorization header was allowed")
}
response := &http.Response{Body: ioNopCloser{Reader: strings.NewReader("12345")}}
if _, err := ReadLimitedBody(response, 4); err == nil {
t.Fatal("oversized response was allowed")
}
store := &audit.MemoryStore{}
if err := RecordPolicyChange(context.Background(), store, "user", "policy", "corr", NetworkPolicy{Revision: 1}, policy); err != nil {
t.Fatal(err)
}
if len(store.Events) != 1 || store.Events[0].Action != "probe.network_policy.update" || store.Events[0].After["maxResponseBytes"] == nil {
t.Fatalf("audit=%+v", store.Events)
}
if err := RecordPolicyChange(context.Background(), nil, "user", "policy", "corr", policy, policy); err == nil {
t.Fatal("nil audit store was accepted")
}
}
type ioNopCloser struct{ *strings.Reader }
func (ioNopCloser) Close() error { return nil }
func resolverHost(host string) string { return host }
+52
View File
@@ -0,0 +1,52 @@
package probe
import (
"encoding/json"
"os"
"path/filepath"
"testing"
)
type executorScenario struct {
SchemaVersion int `json:"schemaVersion"`
ID string `json:"id"`
InitialState struct {
Cases []struct {
ID string `json:"id"`
Type string `json:"type"`
ExpectedState string `json:"expectedState"`
ErrorClass string `json:"errorClass"`
} `json:"cases"`
CredentialExpectation string `json:"credentialExpectation"`
} `json:"initialState"`
ExpectedOutcomes []struct {
Assertion string `json:"assertion"`
} `json:"expectedOutcomes"`
}
func TestProbeExecutorScenarioFixtureCoversAcceptanceCases(t *testing.T) {
fixturePath := filepath.Join("..", "..", "fixtures", "scenarios", "probe-executor-cases.json")
content, err := os.ReadFile(fixturePath)
if err != nil {
t.Fatal(err)
}
var scenario executorScenario
if err := json.Unmarshal(content, &scenario); err != nil {
t.Fatal(err)
}
if scenario.SchemaVersion != 1 || scenario.ID != "probe-executor-cases" || len(scenario.InitialState.Cases) < 9 || len(scenario.ExpectedOutcomes) < 3 {
t.Fatalf("fixture is incomplete: %+v", scenario)
}
seen := make(map[string]bool)
for _, testCase := range scenario.InitialState.Cases {
seen[testCase.Type+":"+testCase.ExpectedState] = true
}
for _, required := range []string{"http:up", "http:down", "http:unknown", "tls:up", "dns:up", "tcp:up", "icmp:unknown"} {
if !seen[required] {
t.Fatalf("fixture is missing acceptance case %q", required)
}
}
if scenario.InitialState.CredentialExpectation == "" {
t.Fatal("fixture does not state the credential redaction expectation")
}
}
+308
View File
@@ -0,0 +1,308 @@
package probe
import (
"context"
"errors"
"fmt"
"sort"
"strings"
"sync"
"time"
)
type Executor interface {
Execute(context.Context, Definition) (Result, error)
}
type RetryableError struct{ Err error }
func (e RetryableError) Error() string {
if e.Err == nil {
return "retryable probe error"
}
return e.Err.Error()
}
func (e RetryableError) Unwrap() error { return e.Err }
type SchedulerConfig struct {
MaxConcurrent int
MaxAttempts int
AttemptTimeout time.Duration
RetryBackoff time.Duration
Now func() time.Time
}
func (c SchedulerConfig) withDefaults() SchedulerConfig {
if c.MaxConcurrent == 0 {
c.MaxConcurrent = 16
}
if c.MaxAttempts == 0 {
c.MaxAttempts = 2
}
if c.AttemptTimeout == 0 {
c.AttemptTimeout = 10 * time.Second
}
if c.RetryBackoff == 0 {
c.RetryBackoff = 100 * time.Millisecond
}
if c.Now == nil {
c.Now = func() time.Time { return time.Now().UTC() }
}
return c
}
func (c SchedulerConfig) Validate() error {
if c.MaxConcurrent < 1 || c.MaxConcurrent > 64 || c.MaxAttempts < 1 || c.MaxAttempts > 3 || c.AttemptTimeout <= 0 || c.AttemptTimeout > 2*time.Minute || c.RetryBackoff < 0 || c.RetryBackoff > time.Minute {
return errors.New("probe scheduler configuration is outside safe bounds")
}
return nil
}
type Metrics struct {
Runs int64 `json:"runs"`
Completed int64 `json:"completed"`
Failed int64 `json:"failed"`
TimedOut int64 `json:"timedOut"`
Retried int64 `json:"retried"`
SkippedOverlap int64 `json:"skippedOverlap"`
Active int64 `json:"active"`
LastRunAt time.Time `json:"lastRunAt"`
}
type RunReport struct {
Results []Result `json:"results"`
Metrics Metrics `json:"metrics"`
}
type Scheduler struct {
executor Executor
config SchedulerConfig
sem chan struct{}
mu sync.Mutex
inflight map[string]context.CancelFunc
closed bool
wg sync.WaitGroup
metrics Metrics
}
func NewScheduler(executor Executor, config SchedulerConfig) (*Scheduler, error) {
config = config.withDefaults()
if executor == nil {
return nil, errors.New("probe executor is required")
}
if err := config.Validate(); err != nil {
return nil, err
}
return &Scheduler{executor: executor, config: config, sem: make(chan struct{}, config.MaxConcurrent), inflight: make(map[string]context.CancelFunc)}, nil
}
func (s *Scheduler) Run(ctx context.Context, definitions []Definition) (RunReport, error) {
if s == nil {
return RunReport{}, errors.New("probe scheduler is nil")
}
if ctx == nil {
return RunReport{}, errors.New("probe scheduler context is nil")
}
if err := ctx.Err(); err != nil {
return RunReport{}, err
}
items := append([]Definition(nil), definitions...)
sort.SliceStable(items, func(i, j int) bool { return items[i].ID < items[j].ID })
for _, definition := range items {
if !definition.Enabled || definition.ArchivedAt != nil {
continue
}
if err := definition.Validate(); err != nil {
return RunReport{}, fmt.Errorf("validate probe %s: %w", definition.ID, err)
}
}
results := make(chan Result, len(items))
for _, definition := range items {
if definition.Enabled && definition.ArchivedAt == nil {
s.begin(definition, ctx, results)
}
}
s.wg.Wait()
close(results)
report := RunReport{Results: make([]Result, 0, len(results))}
for result := range results {
report.Results = append(report.Results, result)
}
sort.SliceStable(report.Results, func(i, j int) bool { return report.Results[i].ProbeID < report.Results[j].ProbeID })
s.mu.Lock()
report.Metrics = s.metrics
s.mu.Unlock()
return report, nil
}
func (s *Scheduler) begin(definition Definition, parent context.Context, results chan<- Result) bool {
probeID := definition.ID
s.mu.Lock()
if s.closed {
s.mu.Unlock()
return false
}
if _, exists := s.inflight[probeID]; exists {
s.metrics.SkippedOverlap++
s.mu.Unlock()
return false
}
child, cancel := context.WithCancel(parent)
s.inflight[probeID] = cancel
s.wg.Add(1)
s.metrics.Runs++
s.metrics.Active++
s.mu.Unlock()
go func() {
defer s.wg.Done()
defer func() { s.mu.Lock(); delete(s.inflight, probeID); s.metrics.Active--; s.mu.Unlock() }()
select {
case s.sem <- struct{}{}:
case <-child.Done():
result := s.failureResult(probeID, child.Err(), 0)
s.recordFailure(result)
results <- result
return
}
defer func() { <-s.sem }()
result := s.executeDefinition(child, definition)
results <- result
}()
return true
}
func (s *Scheduler) executeDefinition(ctx context.Context, definition Definition) Result {
started := s.config.Now().UTC()
var lastErr error
var attempts int
for attempts = 1; attempts <= s.config.MaxAttempts; attempts++ {
attemptCtx, cancel := context.WithTimeout(ctx, s.config.AttemptTimeout)
result, err := s.executor.Execute(attemptCtx, definition)
deadline := errors.Is(attemptCtx.Err(), context.DeadlineExceeded)
cancel()
if err == nil {
result = normalizeResult(result, definition.ID, attempts, started, s.config.Now)
s.recordSuccess(result, false)
return result
}
lastErr = err
if errors.Is(ctx.Err(), context.Canceled) || errors.Is(ctx.Err(), context.DeadlineExceeded) {
break
}
if attempts >= s.config.MaxAttempts || (!deadline && !isRetryable(err)) {
break
}
s.mu.Lock()
s.metrics.Retried++
s.mu.Unlock()
if !waitBackoff(ctx, s.config.RetryBackoff) {
break
}
}
result := s.failureResult(definition.ID, lastErr, attempts)
if errors.Is(ctx.Err(), context.Canceled) || errors.Is(ctx.Err(), context.DeadlineExceeded) {
result.ErrorClass = "canceled"
}
s.recordFailure(result)
return result
}
func (s *Scheduler) failureResult(probeID string, err error, attempts int) Result {
result := Result{ProbeID: probeID, State: "unknown", Attempts: attempts, ErrorClass: "execution_error", ErrorMessage: boundedError(err)}
if errors.Is(err, context.DeadlineExceeded) {
result.ErrorClass = "timeout"
}
if errors.Is(err, context.Canceled) {
result.ErrorClass = "canceled"
}
return normalizeResult(result, probeID, attempts, s.config.Now(), s.config.Now)
}
func (s *Scheduler) recordSuccess(result Result, _ bool) {
s.mu.Lock()
s.metrics.Completed++
s.metrics.LastRunAt = result.CompletedAt
s.mu.Unlock()
}
func (s *Scheduler) recordFailure(result Result) {
s.mu.Lock()
s.metrics.Failed++
if result.ErrorClass == "timeout" {
s.metrics.TimedOut++
}
s.metrics.LastRunAt = result.CompletedAt
s.mu.Unlock()
}
func (s *Scheduler) Shutdown(ctx context.Context) error {
if s == nil {
return nil
}
if ctx == nil {
ctx = context.Background()
}
s.mu.Lock()
if !s.closed {
s.closed = true
for _, cancel := range s.inflight {
cancel()
}
}
s.mu.Unlock()
done := make(chan struct{})
go func() { s.wg.Wait(); close(done) }()
select {
case <-done:
return nil
case <-ctx.Done():
return ctx.Err()
}
}
func (s *Scheduler) Metrics() Metrics { s.mu.Lock(); defer s.mu.Unlock(); return s.metrics }
func normalizeResult(result Result, probeID string, attempts int, started time.Time, now func() time.Time) Result {
result.ProbeID = probeID
if result.State != "up" && result.State != "degraded" && result.State != "down" && result.State != "unknown" {
result.State = "unknown"
if result.ErrorClass == "" {
result.ErrorClass = "invalid_result"
}
}
result.Attempts = attempts
if result.ObservedAt.IsZero() {
result.ObservedAt = started.UTC()
}
if result.CompletedAt.IsZero() {
result.CompletedAt = now().UTC()
}
result.ErrorMessage = boundText(result.ErrorMessage, 256)
return result
}
func isRetryable(err error) bool { var retryable RetryableError; return errors.As(err, &retryable) }
func waitBackoff(ctx context.Context, delay time.Duration) bool {
if delay <= 0 {
return true
}
timer := time.NewTimer(delay)
defer timer.Stop()
select {
case <-timer.C:
return true
case <-ctx.Done():
return false
}
}
func boundedError(err error) string {
if err == nil {
return ""
}
return boundText(err.Error(), 256)
}
func boundText(value string, max int) string {
value = strings.TrimSpace(value)
if len(value) > max {
return value[:max]
}
return value
}
+160
View File
@@ -0,0 +1,160 @@
package probe
import (
"context"
"errors"
"fmt"
"sync"
"testing"
"time"
)
type executorFunc func(context.Context, Definition) (Result, error)
func (f executorFunc) Execute(ctx context.Context, definition Definition) (Result, error) {
return f(ctx, definition)
}
func schedulerDefinition(id string) Definition {
return Definition{ID: id, ServiceID: "service", Name: id, Type: TypeHTTP, Target: Target{Scheme: "https", Host: "example.internal", Port: 443}, Interval: 30 * time.Second, Timeout: 5 * time.Second, Enabled: true, Revision: 1}
}
func TestSchedulerRetriesTimeoutsAndNormalizesResults(t *testing.T) {
var mu sync.Mutex
attempts := 0
executor := executorFunc(func(ctx context.Context, definition Definition) (Result, error) {
if definition.ID != "probe-1" {
t.Errorf("definition=%+v", definition)
}
mu.Lock()
attempts++
current := attempts
mu.Unlock()
if current == 1 {
return Result{}, RetryableError{Err: errors.New("temporary")}
}
return Result{State: "up"}, nil
})
scheduler, err := NewScheduler(executor, SchedulerConfig{MaxConcurrent: 2, MaxAttempts: 2, RetryBackoff: 0})
if err != nil {
t.Fatal(err)
}
report, err := scheduler.Run(context.Background(), []Definition{schedulerDefinition("probe-1")})
if err != nil || len(report.Results) != 1 {
t.Fatalf("report=%+v err=%v", report, err)
}
if report.Results[0].State != "up" || report.Results[0].Attempts != 2 || report.Metrics.Retried != 1 {
t.Fatalf("report=%+v", report)
}
timeoutExecutor := executorFunc(func(ctx context.Context, _ Definition) (Result, error) { <-ctx.Done(); return Result{}, ctx.Err() })
timeoutScheduler, err := NewScheduler(timeoutExecutor, SchedulerConfig{MaxAttempts: 2, AttemptTimeout: 5 * time.Millisecond, RetryBackoff: 0})
if err != nil {
t.Fatal(err)
}
timeoutReport, err := timeoutScheduler.Run(context.Background(), []Definition{schedulerDefinition("timeout")})
if err != nil || timeoutReport.Results[0].ErrorClass != "timeout" || timeoutReport.Results[0].Attempts != 2 || timeoutReport.Metrics.TimedOut != 1 {
t.Fatalf("timeout=%+v err=%v", timeoutReport, err)
}
}
func TestSchedulerPreventsOverlappingRunsAndShutsDown(t *testing.T) {
started := make(chan struct{})
release := make(chan struct{})
executor := executorFunc(func(ctx context.Context, _ Definition) (Result, error) {
closeOnce(started)
select {
case <-release:
return Result{State: "up"}, nil
case <-ctx.Done():
return Result{}, ctx.Err()
}
})
scheduler, err := NewScheduler(executor, SchedulerConfig{MaxAttempts: 1, AttemptTimeout: time.Second})
if err != nil {
t.Fatal(err)
}
first := make(chan RunReport, 1)
go func() {
report, _ := scheduler.Run(context.Background(), []Definition{schedulerDefinition("same")})
first <- report
}()
<-started
second, err := scheduler.Run(context.Background(), []Definition{schedulerDefinition("same")})
if err != nil || len(second.Results) != 0 || second.Metrics.SkippedOverlap < 1 {
t.Fatalf("overlap=%+v err=%v", second, err)
}
shutdownDone := make(chan error, 1)
go func() { shutdownDone <- scheduler.Shutdown(context.Background()) }()
select {
case err := <-shutdownDone:
if err != nil {
t.Fatal(err)
}
case <-time.After(time.Second):
t.Fatal("shutdown did not cancel active probe")
}
close(release)
select {
case <-first:
case <-time.After(time.Second):
t.Fatal("run did not finish after shutdown")
}
if err := scheduler.Shutdown(context.Background()); err != nil {
t.Fatal(err)
}
}
func TestSchedulerTargetScale300IsBoundedAndDeterministic(t *testing.T) {
var mu sync.Mutex
active, maximum := 0, 0
executor := executorFunc(func(ctx context.Context, definition Definition) (Result, error) {
mu.Lock()
active++
if active > maximum {
maximum = active
}
mu.Unlock()
select {
case <-ctx.Done():
return Result{}, ctx.Err()
default:
}
mu.Lock()
active--
mu.Unlock()
return Result{State: "up"}, nil
})
scheduler, err := NewScheduler(executor, SchedulerConfig{MaxConcurrent: 16, MaxAttempts: 1})
if err != nil {
t.Fatal(err)
}
definitions := make([]Definition, 300)
for index := range definitions {
definitions[index] = schedulerDefinition(fmt.Sprintf("probe-%03d", 300-index))
}
started := time.Now()
report, err := scheduler.Run(context.Background(), definitions)
if err != nil || len(report.Results) != 300 {
t.Fatalf("count=%d err=%v", len(report.Results), err)
}
if maximum > 16 {
t.Fatalf("maximum concurrency=%d", maximum)
}
if time.Since(started) > 2*time.Second {
t.Fatalf("300-probe run exceeded budget: %s", time.Since(started))
}
for index := 1; index < len(report.Results); index++ {
if report.Results[index-1].ProbeID > report.Results[index].ProbeID {
t.Fatal("results are not deterministic")
}
}
}
func closeOnce(channel chan struct{}) {
select {
case <-channel:
default:
close(channel)
}
}
+92
View File
@@ -0,0 +1,92 @@
package probe
import (
"errors"
"strings"
"time"
)
const (
TypeHTTP = "http"
TypeTCP = "tcp"
TypeDNS = "dns"
TypeICMP = "icmp"
TypeTLS = "tls"
)
type Target struct {
Scheme string `json:"scheme,omitempty"`
Host string `json:"host"`
Port int `json:"port,omitempty"`
Path string `json:"path,omitempty"`
}
type Definition struct {
ID string `json:"id"`
ServiceID string `json:"serviceId"`
EndpointID string `json:"endpointId,omitempty"`
Name string `json:"name"`
Type string `json:"type"`
Target Target `json:"target"`
Interval time.Duration `json:"interval"`
Timeout time.Duration `json:"timeout"`
Enabled bool `json:"enabled"`
ExpectedStatusCodes []int `json:"expectedStatusCodes,omitempty"`
FollowRedirects bool `json:"followRedirects"`
VerifyTLS bool `json:"verifyTls"`
ContentAssertion map[string]any `json:"contentAssertion,omitempty"`
SecretReference string `json:"secretReference,omitempty"`
NetworkPolicyID string `json:"networkPolicyId,omitempty"`
Revision int64 `json:"revision"`
ArchivedAt *time.Time `json:"archivedAt,omitempty"`
CreatedAt time.Time `json:"createdAt"`
UpdatedAt time.Time `json:"updatedAt"`
}
type Result struct {
ID string `json:"id"`
ProbeID string `json:"probeId"`
ObservedAt time.Time `json:"observedAt"`
CompletedAt time.Time `json:"completedAt"`
State string `json:"state"`
ResponseTimeMS *int `json:"responseTimeMs,omitempty"`
StatusCode *int `json:"statusCode,omitempty"`
Attempts int `json:"attempts"`
ErrorClass string `json:"errorClass,omitempty"`
ErrorMessage string `json:"errorMessage,omitempty"`
Attributes map[string]any `json:"attributes,omitempty"`
Certificate *Certificate `json:"certificate,omitempty"`
}
type Certificate struct {
ID string `json:"id"`
ServiceID string `json:"serviceId"`
EndpointID string `json:"endpointId,omitempty"`
ObservedAt time.Time `json:"observedAt"`
ExpiresAt *time.Time `json:"expiresAt,omitempty"`
Issuer string `json:"issuer,omitempty"`
Subject string `json:"subject,omitempty"`
HostnameValid *bool `json:"hostnameValid,omitempty"`
VerificationState string `json:"verificationState"`
}
func (d Definition) Validate() error {
if strings.TrimSpace(d.ID) == "" || strings.TrimSpace(d.ServiceID) == "" || strings.TrimSpace(d.Name) == "" || len(d.Name) > 160 || strings.TrimSpace(d.Target.Host) == "" || len(d.Target.Host) > 253 || d.Revision < 1 {
return errors.New("probe identity is invalid")
}
if d.Type != TypeHTTP && d.Type != TypeTCP && d.Type != TypeDNS && d.Type != TypeICMP && d.Type != TypeTLS {
return errors.New("probe type is invalid")
}
if d.Interval < 5*time.Second || d.Interval > 24*time.Hour || d.Timeout < time.Second || d.Timeout > 120*time.Second || d.Timeout >= d.Interval {
return errors.New("probe timing is invalid")
}
if d.Target.Port < 0 || d.Target.Port > 65535 || len(d.Target.Path) > 2048 {
return errors.New("probe target is invalid")
}
for _, code := range d.ExpectedStatusCodes {
if code < 100 || code > 599 {
return errors.New("probe expected status is invalid")
}
}
return nil
}
+19
View File
@@ -0,0 +1,19 @@
package probe
import "testing"
func TestDefinitionValidationAndHistoryShape(t *testing.T) {
definition := Definition{ID: "probe", ServiceID: "service", Name: "HTTPS", Type: TypeHTTP, Target: Target{Scheme: "https", Host: "example.internal", Port: 443, Path: "/health"}, Interval: 30e9, Timeout: 5e9, Enabled: true, Revision: 1, ExpectedStatusCodes: []int{200, 204}}
if err := definition.Validate(); err != nil {
t.Fatal(err)
}
definition.Timeout = definition.Interval
if err := definition.Validate(); err == nil {
t.Fatal("expected timeout/interval validation")
}
definition.Timeout = 5e9
definition.ExpectedStatusCodes = []int{700}
if err := definition.Validate(); err == nil {
t.Fatal("expected status validation")
}
}