This commit is contained in:
@@ -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 }
|
||||
@@ -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) }
|
||||
@@ -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
|
||||
}
|
||||
@@ -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 }
|
||||
@@ -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")
|
||||
}
|
||||
}
|
||||
@@ -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
|
||||
}
|
||||
@@ -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)
|
||||
}
|
||||
}
|
||||
@@ -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
|
||||
}
|
||||
@@ -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")
|
||||
}
|
||||
}
|
||||
Reference in New Issue
Block a user