From 52828f84f5e21eaed776454c1bfb603fe4b66b98 Mon Sep 17 00:00:00 2001 From: Benjamin Date: Wed, 30 Sep 2026 16:19:50 +0200 Subject: [PATCH 01/54] fix(install): restart on update, keep the env file, offline install, custom data paths - Replace the binary through a temporary file and a rename, then restart the service when it already runs (start it otherwise). - Merge the env file key by key: only the keys passed as flags change, every other line (comments, DOCKER_HOST, names with digits) is kept. - Take the flag lists from the binary: every FlagTypeBool flag works bare or as --flag=true|false, --flag=value works everywhere, unknown flags are refused. The help lists every flag; a Go test keeps both in step. - Add --binary and --sha256sums for an install with no network access. - Follow --data-dir and --db: create those directories for the service user and render the unit with them as working and writable paths. - Show the configured listen address in the summary, not the script's own environment. - Exit 20 on a failed download and 30 on a filesystem error. - Keep the telemetry directory next to the database instead of /data/shm, which the native unit cannot write. --- deploy/install/install.sh | 497 +++++++++++++++++++------ deploy/install/maintenant.service | 5 +- deploy/install/test/install_args.bats | 178 +++++++++ deploy/install/test/install_basic.bats | 14 + deploy/install/test/offline.bats | 101 +++++ deploy/install/test/service.bats | 243 ++++++++++++ deploy/install/test/uninstall.bats | 31 ++ internal/app/app.go | 1 + internal/app/app_test.go | 28 ++ internal/app/install_script_test.go | 92 +++++ internal/app/main_test.go | 4 + internal/telemetry/telemetry.go | 7 +- 12 files changed, 1071 insertions(+), 130 deletions(-) create mode 100644 deploy/install/test/offline.bats create mode 100644 deploy/install/test/service.bats create mode 100644 internal/app/install_script_test.go diff --git a/deploy/install/install.sh b/deploy/install/install.sh index 844d49fb..0b5326bd 100755 --- a/deploy/install/install.sh +++ b/deploy/install/install.sh @@ -18,6 +18,21 @@ SERVICE_FILE="${SERVICE_FILE:-/etc/systemd/system/maintenant.service}" GITHUB_REPO="kOlapsis/maintenant" GITHUB_API="https://api.github.com" SCRIPT_VERSION="__GIT_SHA__" +NL=' +' + +# Kept in step with internal/app/flags.go by internal/app/install_script_test.go. +BOOL_FLAGS="proxyLabels disableOsEolRefresh disableTelemetry allowPrivateWebhooks \ +mcp mcpAllowUnauthenticated grpc-tls-insecure grpc-insecure-skip-tls-verify embedded-agent" +VALUE_FLAGS="addr baseUrl corsOrigins trustedProxies db organisationName runtime logLevel \ +maxBodySize updateInterval securityScoreThreshold licenseKey \ +smtpHost smtpPort smtpUsername smtpPassword smtpFrom \ +mcpClientId mcpClientSecret mcpAllowedRedirectUris k8sNamespaces k8sExcludeNamespaces \ +statusUrl containerDownAfter retentionSnapshots retentionInterval retentionBatchSize \ +mode server enrollment-token label nodeName grpc-listen grpc-url grpc-tls-cert grpc-tls-key \ +agentRateLimitPerSecond agentStaleThresholdSeconds \ +agentSpoolMaxMemoryBytes agentSpoolMaxDiskBytes agentSpoolMaxAgeSeconds \ +data-dir ca-cert database-url" # ── Color / output ──────────────────────────────────────────────────────────── @@ -38,11 +53,17 @@ abort() { exit "${2:-1}" } +_fs() { + "$@" || abort "Filesystem operation failed: $*" 30 +} + # ── Cleanup trap ────────────────────────────────────────────────────────────── TMPDIR_INSTALL="${TMPDIR_INSTALL:-}" +BINARY_TMP="" cleanup() { if [ -n "${TMPDIR_INSTALL:-}" ]; then rm -rf "$TMPDIR_INSTALL"; fi + if [ -n "${BINARY_TMP:-}" ]; then rm -f "$BINARY_TMP"; fi } trap cleanup EXIT @@ -57,18 +78,31 @@ Script flags: --uninstall Remove Maintenant (keeps data and user by default) --purge With --uninstall: also remove data dir, config, user --skip-cosign Skip cosign signature check (SHA256 still required) + --binary Install this local binary instead of downloading one + (offline install: no network access at all) + --sha256sums With --binary: check it against this SHA256SUMS file --help, -h Show this help -Binary configuration flags (written to /etc/maintenant/maintenant.env): +Binary configuration flags (written to /etc/maintenant/maintenant.env). +Value flags take "--flag value" or "--flag=value". Boolean flags take "--flag" +or "--flag=true|false". Run "maintenant --help" for what each flag does. --addr --baseUrl + --corsOrigins + --trustedProxies --db + --containerDownAfter + --retentionSnapshots + --retentionInterval + --retentionBatchSize --organisationName - --corsOrigins + --statusUrl --runtime + --proxyLabels --logLevel --maxBodySize --updateInterval + --disableOsEolRefresh --securityScoreThreshold --disableTelemetry --allowPrivateWebhooks @@ -85,14 +119,11 @@ Binary configuration flags (written to /etc/maintenant/maintenant.env): --mcpAllowUnauthenticated --k8sNamespaces --k8sExcludeNamespaces - --statusUrl - --retentionSnapshots - --retentionInterval - --retentionBatchSize --mode --server --enrollment-token --label + --nodeName --grpc-listen --grpc-url --grpc-tls-cert @@ -101,9 +132,12 @@ Binary configuration flags (written to /etc/maintenant/maintenant.env): --grpc-insecure-skip-tls-verify --agentRateLimitPerSecond --agentStaleThresholdSeconds - --data-dir + --agentSpoolMaxMemoryBytes + --agentSpoolMaxDiskBytes + --agentSpoolMaxAgeSeconds --embedded-agent --ca-cert + --data-dir --database-url Examples: @@ -114,8 +148,11 @@ Examples: install.sh --mode agent --server grpcs://maintenant.example.com:8443 \ --enrollment-token TOKEN --label web-01 + # Offline, from a binary and its SHA256SUMS copied onto this host + install.sh --binary ./maintenant-v1.2.3-linux-amd64 --sha256sums ./SHA256SUMS + Environment variables: - MAINTENANT_VERSION Version to install (default: latest) + MAINTENANT_VERSION Version to download (default: latest) MAINTENANT_INSTALL_DIR Binary install path (default: /usr/local/bin) MAINTENANT_DATA_DIR Data directory (default: /var/lib/maintenant) MAINTENANT_CONFIG_DIR Config directory (default: /etc/maintenant) @@ -142,9 +179,8 @@ detect_platform() { check_prereqs() { [ "$(id -u)" -eq 0 ] || abort "This script must be run as root (EUID 0)" 11 - # curl or wget - if ! command -v curl >/dev/null 2>&1 && ! command -v wget >/dev/null 2>&1; then - abort "curl or wget is required" 12 + if [ -z "${LOCAL_BINARY:-}" ] && ! command -v curl >/dev/null 2>&1 && ! command -v wget >/dev/null 2>&1; then + abort "curl or wget is required (or install a local binary with --binary)" 12 fi # No tar: the release assets are bare binaries, not archives. @@ -250,8 +286,10 @@ download_and_verify() { BASE_URL="https://github.com/$GITHUB_REPO/releases/download/${VERSION}" log_step "Downloading $ASSET_NAME..." - fetch_url_to "$BASE_URL/$ASSET_NAME" "$TMPDIR_INSTALL/$ASSET_NAME" - fetch_url_to "$BASE_URL/SHA256SUMS" "$TMPDIR_INSTALL/SHA256SUMS" + fetch_url_to "$BASE_URL/$ASSET_NAME" "$TMPDIR_INSTALL/$ASSET_NAME" \ + || abort "Failed to download $ASSET_NAME" 20 + fetch_url_to "$BASE_URL/SHA256SUMS" "$TMPDIR_INSTALL/SHA256SUMS" \ + || abort "Failed to download SHA256SUMS" 20 log_step "Verifying SHA256 checksum..." (cd "$TMPDIR_INSTALL" && sha256sum -c SHA256SUMS --ignore-missing) \ @@ -267,7 +305,8 @@ download_and_verify() { # the signature either. log_warn "cosign 3 or later is required to read the release bundle — skipping signature verification" else - fetch_url_to "$BASE_URL/SHA256SUMS.bundle" "$TMPDIR_INSTALL/SHA256SUMS.bundle" + fetch_url_to "$BASE_URL/SHA256SUMS.bundle" "$TMPDIR_INSTALL/SHA256SUMS.bundle" \ + || abort "Failed to download SHA256SUMS.bundle" 20 log_step "Verifying cosign signature..." if ! cosign verify-blob \ --bundle "$TMPDIR_INSTALL/SHA256SUMS.bundle" \ @@ -278,6 +317,42 @@ download_and_verify() { fi log_info "cosign signature verified" fi + + BINARY_SRC="$TMPDIR_INSTALL/$ASSET_NAME" +} + +# ── use_local_binary ────────────────────────────────────────────────────────── + +use_local_binary() { + [ -f "$LOCAL_BINARY" ] || abort "Binary not found: $LOCAL_BINARY" 2 + TMPDIR_INSTALL=$(mktemp -d) + BINARY_SRC="$LOCAL_BINARY" + VERSION="local" + + if [ -z "${LOCAL_SUMS:-}" ]; then + log_warn "No --sha256sums given: the integrity of $LOCAL_BINARY is not checked" + return + fi + [ -f "$LOCAL_SUMS" ] || abort "SHA256SUMS file not found: $LOCAL_SUMS" 2 + + log_step "Verifying SHA256 checksum against $LOCAL_SUMS..." + LOCAL_SUM=$(sha256sum "$LOCAL_BINARY" | cut -d ' ' -f 1) + MATCHED_ASSET=$(awk -v sum="$LOCAL_SUM" -v suffix="-linux-$ARCH" ' + $1 == sum { + name = $2 + sub(/^\*/, "", name) + if (name ~ /^maintenant-/ && substr(name, length(name) - length(suffix) + 1) == suffix) { + print name + exit + } + }' "$LOCAL_SUMS") + [ -n "$MATCHED_ASSET" ] \ + || abort "SHA256 of $LOCAL_BINARY matches no linux-$ARCH binary listed in $LOCAL_SUMS" 21 + + VERSION="${MATCHED_ASSET#maintenant-}" + VERSION="${VERSION%-linux-"$ARCH"}" + log_info "Checksum matches $MATCHED_ASSET" + log_warn "Offline install: the cosign signature of $LOCAL_SUMS is not verified" } # ── ensure_user ─────────────────────────────────────────────────────────────── @@ -288,7 +363,8 @@ ensure_user() { else log_step "Creating system user $SERVICE_USER..." useradd -r -s /usr/sbin/nologin -d "$DATA_DIR" \ - -c "Maintenant service user" "$SERVICE_USER" + -c "Maintenant service user" "$SERVICE_USER" \ + || abort "Failed to create user $SERVICE_USER" log_info "User $SERVICE_USER created" fi @@ -302,31 +378,126 @@ ensure_user() { fi } +# ── resolve_paths ───────────────────────────────────────────────────────────── +# A flag given now wins over the env file, which wins over the defaults. + +_env_file_value() { + [ -f "$2" ] || return 0 + awk -v key="$1" -v q="'" ' + { + line = $0 + sub(/^[ \t]+/, "", line) + if (!match(line, /^[A-Za-z_][A-Za-z0-9_]*[ \t]*=/)) next + k = substr(line, 1, RLENGTH - 1) + sub(/[ \t]+$/, "", k) + if (k != key) next + v = substr(line, RLENGTH + 1) + sub(/^[ \t]+/, "", v) + sub(/[ \t]+$/, "", v) + if (length(v) >= 2 && (v ~ /^".*"$/ || v ~ ("^" q ".*" q "$"))) v = substr(v, 2, length(v) - 2) + found = v + } + END { printf "%s", found } + ' "$2" +} + +_strip_trailing_slashes() { + p="$1" + while [ "$p" != "/" ]; do + case "$p" in + */) p="${p%/}" ;; + *) break ;; + esac + done + printf '%s' "$p" +} + +resolve_paths() { + ENV_FILE="$CONFIG_DIR/maintenant.env" + + if _has_binary_flag data-dir; then + DATA_DIR=$(_binary_flag_value data-dir) + else + from_file=$(_env_file_value MAINTENANT_DATA_DIR "$ENV_FILE") + [ -z "$from_file" ] || DATA_DIR="$from_file" + fi + DATA_DIR=$(_strip_trailing_slashes "$DATA_DIR") + case "$DATA_DIR" in + /?*) ;; + *) abort "The data directory must be an absolute path other than /: $DATA_DIR" 2 ;; + esac + + if _has_binary_flag db; then + DB_PATH=$(_binary_flag_value db) + else + DB_PATH=$(_env_file_value MAINTENANT_DB "$ENV_FILE") + fi + [ -n "$DB_PATH" ] || DB_PATH="./maintenant.db" + while :; do + case "$DB_PATH" in + ./*) DB_PATH="${DB_PATH#./}" ;; + *) break ;; + esac + done + case "$DB_PATH" in + /*) ;; + *) DB_PATH="$DATA_DIR/$DB_PATH" ;; + esac + DB_DIR=$(dirname "$DB_PATH") + + for unit_path in "$INSTALL_DIR" "$CONFIG_DIR" "$DATA_DIR" "$DB_DIR"; do + case "$unit_path" in + *[!A-Za-z0-9._/@+-]*) + abort "Unsupported character in path (letters, digits and ._/@+- only): $unit_path" 2 ;; + esac + done +} + +# ── prepare_dirs ────────────────────────────────────────────────────────────── + +_refuse_foreign_dir() { + [ -d "$1" ] || return 0 + [ -z "$(find "$1" -prune -user "$SERVICE_USER" 2>/dev/null)" ] || return 0 + [ -z "$(ls -A "$1")" ] && return 0 + abort "$1 already exists, is not empty and does not belong to $SERVICE_USER: use a dedicated directory, or chown it to $SERVICE_USER first" 30 +} + +_own_dir() { + _fs mkdir -p "$1" + _fs chown "$SERVICE_USER:$SERVICE_USER" "$1" + _fs chmod 0750 "$1" +} + +prepare_dirs() { + _refuse_foreign_dir "$DATA_DIR" + _refuse_foreign_dir "$DB_DIR" + + _fs mkdir -p "$CONFIG_DIR" + _fs chown "root:$SERVICE_USER" "$CONFIG_DIR" + _fs chmod 0750 "$CONFIG_DIR" + + _own_dir "$DATA_DIR" + [ "$DB_DIR" = "$DATA_DIR" ] || _own_dir "$DB_DIR" +} + # ── install_binary ──────────────────────────────────────────────────────────── install_binary() { log_step "Installing binary..." - mkdir -p "$DATA_DIR" - chown "$SERVICE_USER:$SERVICE_USER" "$DATA_DIR" - chmod 0750 "$DATA_DIR" - - mkdir -p "$CONFIG_DIR" - chown "root:$SERVICE_USER" "$CONFIG_DIR" - chmod 0750 "$CONFIG_DIR" - - install -m 0755 -o root -g root \ - "$TMPDIR_INSTALL/maintenant-${VERSION}-linux-${ARCH}" \ - "$INSTALL_DIR/maintenant" + BINARY_TMP=$(mktemp "$INSTALL_DIR/.maintenant.XXXXXX") \ + || abort "Cannot create a temporary file in $INSTALL_DIR" 30 + _fs install -m 0755 -o root -g root "$BINARY_SRC" "$BINARY_TMP" + _fs mv -f "$BINARY_TMP" "$INSTALL_DIR/maintenant" + BINARY_TMP="" log_info "Binary installed to $INSTALL_DIR/maintenant" } # ── install_service ─────────────────────────────────────────────────────────── -install_service() { - [ -z "${NO_SERVICE:-}" ] || { log_info "Skipping service installation (--no-service)"; return; } - - log_step "Installing systemd service..." - cat > "$SERVICE_FILE" <<'UNIT' +render_unit() { + rw_paths="$DATA_DIR" + [ "$DB_DIR" = "$DATA_DIR" ] || rw_paths="$rw_paths $DB_DIR" + cat </dev/null 2>&1 && systemctl is-active --quiet maintenant 2>/dev/null; then + log_warn "maintenant.service still runs the previous binary: systemctl restart maintenant" + fi + return + fi + + log_step "Installing systemd service..." + render_unit > "$SERVICE_FILE" || abort "Failed to write $SERVICE_FILE" 30 + + systemctl daemon-reload || abort "systemctl daemon-reload failed" 31 + systemctl enable maintenant || abort "systemctl enable maintenant failed" 31 + if systemctl is-active --quiet maintenant; then + log_step "Restarting the service on the new binary..." + systemctl restart maintenant || abort "systemctl restart maintenant failed" 31 + else + log_step "Starting the service..." + systemctl start maintenant || abort "systemctl start maintenant failed" 31 + fi log_step "Waiting for service to become active..." i=0 @@ -381,7 +573,12 @@ UNIT # ── print_summary ───────────────────────────────────────────────────────────── print_summary() { - LISTEN_ADDR="${MAINTENANT_ADDR:-127.0.0.1:8080}" + if _has_binary_flag addr; then + LISTEN_ADDR=$(_binary_flag_value addr) + else + LISTEN_ADDR=$(_env_file_value MAINTENANT_ADDR "$CONFIG_DIR/maintenant.env") + fi + [ -n "$LISTEN_ADDR" ] || LISTEN_ADDR="127.0.0.1:8080" cat < "$TMPDIR_INSTALL/new_flags.env" if [ ! -f "$ENV_FILE" ]; then - # First creation { printf '%s\n\n' "$HEADER" cat "$TMPDIR_INSTALL/new_flags.env" - } > "$ENV_FILE" - chown "root:$SERVICE_USER" "$ENV_FILE" - chmod 0640 "$ENV_FILE" + } > "$TMPDIR_INSTALL/merged.env" || abort "Cannot write $TMPDIR_INSTALL/merged.env" 30 + _install_env_file log_info "Created $ENV_FILE" return fi - # Merge: read existing, override with new - KEYS_COUNT=0 - UPDATES=0 - - # Read existing non-comment lines - grep -E '^MAINTENANT_[A-Z_]+=.*$' "$ENV_FILE" > "$TMPDIR_INSTALL/existing.env" 2>/dev/null || true - - # Build merged file { printf '%s\n\n' "$HEADER" + awk -v counts="$TMPDIR_INSTALL/merge.counts" ' + BEGIN { head = 1 } + FILENAME == ARGV[1] { + i = index($0, "=") + if (i > 1) { + k = substr($0, 1, i - 1) + if (!(k in val)) order[++n] = k + val[k] = substr($0, i + 1) + } + next + } + head && /^# (Generated by install\.sh$|Last updated: |Script version: |Edit this file directly then: )/ { generated = 1; next } + head && generated && /^[ \t]*$/ { head = 0; next } + { head = 0 } + { + line = $0 + sub(/^[ \t]+/, "", line) + if (match(line, /^[A-Za-z_][A-Za-z0-9_]*[ \t]*=/)) { + k = substr(line, 1, RLENGTH - 1) + sub(/[ \t]+$/, "", k) + if (k in val) { + print k "=" val[k] + written[k] = 1 + updated++ + next + } + } + print + } + END { + for (j = 1; j <= n; j++) { + if (!(order[j] in written)) { + print order[j] "=" val[order[j]] + added++ + } + } + print (updated + 0) " " (added + 0) > counts + } + ' "$TMPDIR_INSTALL/new_flags.env" "$ENV_FILE" + } > "$TMPDIR_INSTALL/merged.env" || abort "Cannot merge into $ENV_FILE" 30 + + read -r UPDATED ADDED < "$TMPDIR_INSTALL/merge.counts" + _install_env_file + log_info "$ENV_FILE updated (${UPDATED} keys replaced, ${ADDED} keys added, other lines kept)" +} - # Start with existing entries, override if in new_flags - while IFS='=' read -r key rest; do - [ -n "$key" ] || continue - val="$rest" - # Check if this key appears in new flags - new_val=$(grep "^${key}=" "$TMPDIR_INSTALL/new_flags.env" | cut -d= -f2-) - if [ -n "$new_val" ]; then - printf '%s=%s\n' "$key" "$new_val" - UPDATES=$((UPDATES + 1)) - else - printf '%s=%s\n' "$key" "$val" - fi - KEYS_COUNT=$((KEYS_COUNT + 1)) - done < "$TMPDIR_INSTALL/existing.env" - - # Add new keys not already in existing - while IFS='=' read -r key rest; do - [ -n "$key" ] || continue - if ! grep -q "^${key}=" "$TMPDIR_INSTALL/existing.env" 2>/dev/null; then - printf '%s=%s\n' "$key" "$rest" - KEYS_COUNT=$((KEYS_COUNT + 1)) - fi - done < "$TMPDIR_INSTALL/new_flags.env" - } > "$TMPDIR_INSTALL/merged.env" - - mv "$TMPDIR_INSTALL/merged.env" "$ENV_FILE" - chown "root:$SERVICE_USER" "$ENV_FILE" - chmod 0640 "$ENV_FILE" - log_info "$ENV_FILE updated (${KEYS_COUNT} keys preserved, ${UPDATES} keys updated)" +_install_env_file() { + _fs chown "root:$SERVICE_USER" "$TMPDIR_INSTALL/merged.env" + _fs chmod 0640 "$TMPDIR_INSTALL/merged.env" + _fs mv -f "$TMPDIR_INSTALL/merged.env" "$ENV_FILE" } # ── uninstall ───────────────────────────────────────────────────────────────── uninstall() { log_step "Uninstalling Maintenant..." + resolve_paths # Stop and disable service systemctl stop maintenant 2>/dev/null || true @@ -623,6 +868,10 @@ _purge() { rm -rf "$DATA_DIR" log_info "Removed $DATA_DIR" fi + case "$DB_DIR/" in + "$DATA_DIR"/*) ;; + *) [ ! -d "$DB_DIR" ] || log_warn "Database directory kept, it lies outside $DATA_DIR: $DB_DIR" ;; + esac if [ -d "$CONFIG_DIR" ]; then rm -rf "$CONFIG_DIR" log_info "Removed $CONFIG_DIR" @@ -670,10 +919,15 @@ main() { log_step "Starting Maintenant installation" detect_platform check_prereqs - resolve_version - download_and_verify + resolve_paths + if [ -n "$LOCAL_BINARY" ]; then + use_local_binary + else + resolve_version + download_and_verify + fi ensure_user - install_binary + prepare_dirs # Apply binary flags to env file if any were provided if [ -n "$BINARY_FLAG_KEYS" ]; then @@ -681,6 +935,7 @@ main() { merge_env_file fi + install_binary install_service print_summary } diff --git a/deploy/install/maintenant.service b/deploy/install/maintenant.service index dca0d3bd..dfbe5063 100644 --- a/deploy/install/maintenant.service +++ b/deploy/install/maintenant.service @@ -1,5 +1,5 @@ [Unit] -Description=Maintenant monitoring +Description=Maintenant infrastructure monitoring Documentation=https://docs.maintenant.dev After=network-online.target Wants=network-online.target @@ -8,14 +8,13 @@ Wants=network-online.target Type=simple User=maintenant Group=maintenant +Environment=MAINTENANT_DATA_DIR=/var/lib/maintenant EnvironmentFile=-/etc/maintenant/maintenant.env ExecStart=/usr/local/bin/maintenant WorkingDirectory=/var/lib/maintenant Restart=on-failure RestartSec=5s LimitNOFILE=65536 - -# Hardening NoNewPrivileges=true PrivateTmp=true ProtectSystem=strict diff --git a/deploy/install/test/install_args.bats b/deploy/install/test/install_args.bats index 23acc5b0..fe713106 100644 --- a/deploy/install/test/install_args.bats +++ b/deploy/install/test/install_args.bats @@ -225,3 +225,181 @@ SCRIPT="$(cd "$(dirname "$BATS_TEST_FILENAME")/.." && pwd)/install.sh" [ "$PERMS" = "640" ] rm -rf "$FAKE_TMPDIR" "$FAKE_CONFIG_DIR" } + +@test "merge_env_file: keeps every line it was not asked to change" { + FAKE_TMPDIR=$(mktemp -d) + FAKE_CONFIG_DIR=$(mktemp -d) + cat > "$FAKE_CONFIG_DIR/maintenant.env" <<'ENV' +# Generated by install.sh +# Last updated: 2020-01-01T00:00:00Z +# Script version: abc1234 +# Edit this file directly then: systemctl restart maintenant + +MAINTENANT_ADDR=127.0.0.1:8080 +# socket proxy in front of the Docker API +DOCKER_HOST=tcp://socket-proxy:2375 +MAINTENANT_K8S_NAMESPACES=prod,staging +ENV + + run bash -c " + NO_COLOR=1 + TMPDIR_INSTALL='$FAKE_TMPDIR' + CONFIG_DIR='$FAKE_CONFIG_DIR' + SERVICE_USER=nobody + chown() { return 0; } + export -f chown + export NO_COLOR TMPDIR_INSTALL CONFIG_DIR SERVICE_USER + _INSTALL_SH_TESTING=1 . '$SCRIPT' + parse_maintenant_flags --addr 0.0.0.0:9000 --logLevel debug + merge_env_file + " + [ "$status" -eq 0 ] + ENV_FILE="$FAKE_CONFIG_DIR/maintenant.env" + grep -qx 'MAINTENANT_ADDR=0.0.0.0:9000' "$ENV_FILE" + grep -qx 'MAINTENANT_LOG_LEVEL=debug' "$ENV_FILE" + grep -qx 'DOCKER_HOST=tcp://socket-proxy:2375' "$ENV_FILE" + grep -qx 'MAINTENANT_K8S_NAMESPACES=prod,staging' "$ENV_FILE" + grep -qx '# socket proxy in front of the Docker API' "$ENV_FILE" + [ "$(grep -c '^MAINTENANT_ADDR=' "$ENV_FILE")" -eq 1 ] + [ "$(grep -c '^# Generated by install.sh$' "$ENV_FILE")" -eq 1 ] + [ -z "$(grep '2020-01-01' "$ENV_FILE")" ] + rm -rf "$FAKE_TMPDIR" "$FAKE_CONFIG_DIR" +} + +@test "merge_env_file: a flag given twice keeps its last value" { + FAKE_TMPDIR=$(mktemp -d) + FAKE_CONFIG_DIR=$(mktemp -d) + printf 'MAINTENANT_ADDR=127.0.0.1:8080\n' > "$FAKE_CONFIG_DIR/maintenant.env" + + run bash -c " + NO_COLOR=1 + TMPDIR_INSTALL='$FAKE_TMPDIR' + CONFIG_DIR='$FAKE_CONFIG_DIR' + SERVICE_USER=nobody + chown() { return 0; } + export -f chown + export NO_COLOR TMPDIR_INSTALL CONFIG_DIR SERVICE_USER + _INSTALL_SH_TESTING=1 . '$SCRIPT' + parse_maintenant_flags --addr 0.0.0.0:1 --addr 0.0.0.0:2 + merge_env_file + " + [ "$status" -eq 0 ] + [ "$(grep '^MAINTENANT_ADDR=' "$FAKE_CONFIG_DIR/maintenant.env")" = "MAINTENANT_ADDR=0.0.0.0:2" ] + rm -rf "$FAKE_TMPDIR" "$FAKE_CONFIG_DIR" +} + +# ── boolean and value flags ─────────────────────────────────────────────────── + +parse_and_print() { + bash -c " + _INSTALL_SH_TESTING=1 NO_COLOR=1 . '$SCRIPT' + parse_maintenant_flags $* + printf '%s' \"\$BINARY_FLAG_KEYS\" | while IFS= read -r k; do + printf '%s=%s\n' \"\$k\" \"\$(_binary_flag_value \"\$k\")\" + done + " +} + +@test "parse_maintenant_flags: a boolean flag does not swallow the next flag" { + run parse_and_print --disableOsEolRefresh --addr 0.0.0.0:8080 + [ "$status" -eq 0 ] + [[ "$output" == *"disableOsEolRefresh=true"* ]] + [[ "$output" == *"addr=0.0.0.0:8080"* ]] +} + +@test "parse_maintenant_flags: a boolean flag may come last" { + run parse_and_print --addr 0.0.0.0:8080 --proxyLabels + [ "$status" -eq 0 ] + [[ "$output" == *"proxyLabels=true"* ]] +} + +@test "parse_maintenant_flags: every boolean flag of the binary works bare" { + run bash -c " + _INSTALL_SH_TESTING=1 NO_COLOR=1 . '$SCRIPT' + for f in \$BOOL_FLAGS; do + parse_maintenant_flags --\$f --addr 0.0.0.0:8080 || exit 1 + [ \"\$(_binary_flag_value \$f)\" = true ] || { echo \"--\$f\"; exit 1; } + [ \"\$(_binary_flag_value addr)\" = 0.0.0.0:8080 ] || { echo \"--\$f\"; exit 1; } + done + " + [ "$status" -eq 0 ] +} + +@test "parse_maintenant_flags: a boolean flag takes =false" { + run parse_and_print --disableTelemetry=false --mcp=TRUE + [ "$status" -eq 0 ] + [[ "$output" == *"disableTelemetry=false"* ]] + [[ "$output" == *"mcp=true"* ]] +} + +@test "parse_maintenant_flags: a boolean flag refuses a value that is not a boolean" { + run parse_and_print --mcp=maybe + [ "$status" -eq 2 ] +} + +@test "parse_maintenant_flags: a value flag takes --flag=value" { + run parse_and_print --addr=0.0.0.0:9000 --smtpPassword=--starts-with-dashes + [ "$status" -eq 0 ] + [[ "$output" == *"addr=0.0.0.0:9000"* ]] + [[ "$output" == *"smtpPassword=--starts-with-dashes"* ]] +} + +@test "parse_maintenant_flags: a value flag refuses the next flag as its value" { + run parse_and_print --addr --mcp + [ "$status" -eq 2 ] +} + +@test "parse_maintenant_flags: a flag the binary does not know is refused" { + run parse_and_print --adr 0.0.0.0:8080 + [ "$status" -eq 2 ] + [[ "$output" == *"Unknown argument: --adr"* ]] +} + +@test "parse_maintenant_flags: the binary's action flags are refused" { + run parse_and_print --copy-store-to postgres://db/maintenant + [ "$status" -eq 2 ] + run parse_and_print --yes + [ "$status" -eq 2 ] +} + +@test "parse_maintenant_flags: a value on two lines is refused" { + run bash -c " + _INSTALL_SH_TESTING=1 NO_COLOR=1 . '$SCRIPT' + parse_maintenant_flags --label \"\$(printf 'web\nMAINTENANT_MODE=server')\" + " + [ "$status" -eq 2 ] +} + +@test "parse_maintenant_flags: flags the help used to omit are accepted" { + run parse_and_print --proxyLabels --disableOsEolRefresh --nodeName node-1 \ + --trustedProxies 10.0.0.0/8 --containerDownAfter 5m \ + --agentSpoolMaxMemoryBytes 1048576 --agentSpoolMaxDiskBytes 10485760 \ + --agentSpoolMaxAgeSeconds 3600 + [ "$status" -eq 0 ] + [[ "$output" == *"nodeName=node-1"* ]] + [[ "$output" == *"agentSpoolMaxAgeSeconds=3600"* ]] +} + +@test "--help lists the flags it used to omit" { + run bash "$SCRIPT" --help + [ "$status" -eq 0 ] + for f in proxyLabels disableOsEolRefresh nodeName trustedProxies containerDownAfter \ + agentSpoolMaxMemoryBytes agentSpoolMaxDiskBytes agentSpoolMaxAgeSeconds binary sha256sums; do + [[ "$output" == *"--$f"* ]] || { echo "missing --$f"; return 1; } + done +} + +@test "parse_maintenant_flags: --binary and --sha256sums are script flags" { + run bash -c " + _INSTALL_SH_TESTING=1 NO_COLOR=1 . '$SCRIPT' + parse_maintenant_flags --binary ./maintenant --sha256sums=./SHA256SUMS + echo \"LOCAL_BINARY=\$LOCAL_BINARY LOCAL_SUMS=\$LOCAL_SUMS KEYS=[\$BINARY_FLAG_KEYS]\" + " + [ "$status" -eq 0 ] + [[ "$output" == *"LOCAL_BINARY=./maintenant LOCAL_SUMS=./SHA256SUMS KEYS=[]"* ]] +} + +@test "parse_maintenant_flags: --sha256sums without --binary exits 2" { + run parse_and_print --sha256sums ./SHA256SUMS + [ "$status" -eq 2 ] +} diff --git a/deploy/install/test/install_basic.bats b/deploy/install/test/install_basic.bats index 68dbd958..4c189139 100644 --- a/deploy/install/test/install_basic.bats +++ b/deploy/install/test/install_basic.bats @@ -129,6 +129,20 @@ run_script() { rm -rf "$FAKE_TMPDIR" } +@test "download_and_verify: a failed download exits 20" { + run bash -c " + VERSION='v1.0.0' + ARCH='amd64' + SKIP_COSIGN=1 + NO_COLOR=1 + export VERSION ARCH SKIP_COSIGN NO_COLOR + _INSTALL_SH_TESTING=1 . '$SCRIPT' + fetch_url_to() { return 22; } + download_and_verify + " + [ "$status" -eq 20 ] +} + # ── full flow --no-service ──────────────────────────────────────────────────── @test "full install with --no-service succeeds" { diff --git a/deploy/install/test/offline.bats b/deploy/install/test/offline.bats new file mode 100644 index 00000000..2e50c1bc --- /dev/null +++ b/deploy/install/test/offline.bats @@ -0,0 +1,101 @@ +#!/usr/bin/env bats +# Tests for the offline install from a local binary (--binary, --sha256sums) +bats_require_minimum_version 1.5.0 + +load 'setup' + +SCRIPT="$(cd "$(dirname "$BATS_TEST_FILENAME")/.." && pwd)/install.sh" + +setup() { + WORK=$(mktemp -d) + mkdir -p "$WORK/media" "$WORK/bin" "$WORK/data" "$WORK/etc" + printf 'maintenant binary' > "$WORK/media/maintenant-v1.2.3-linux-amd64" + (cd "$WORK/media" && sha256sum maintenant-v1.2.3-linux-amd64 > SHA256SUMS) +} + +teardown() { + rm -rf "$WORK" +} + +# Every network path fails and leaves a trace in $WORK/network.log. +run_offline() { + local svc + svc=$(id -un) + run bash -c " + NO_COLOR=1 + uname() { case \"\$1\" in -s) echo Linux;; -m) echo x86_64;; esac; } + id() { echo '0'; } + useradd() { return 0; } + getent() { return 1; } + chown() { return 0; } + install() { cp \"\${@: -2:1}\" \"\${@: -1}\"; } + curl() { echo \"curl \$*\" >> '$WORK/network.log'; return 7; } + wget() { echo \"wget \$*\" >> '$WORK/network.log'; return 4; } + export -f uname id useradd getent chown install curl wget + INSTALL_DIR='$WORK/bin' + DATA_DIR='$WORK/data' + CONFIG_DIR='$WORK/etc' + SERVICE_USER='$svc' + export NO_COLOR INSTALL_DIR DATA_DIR CONFIG_DIR + _INSTALL_SH_TESTING=1 . '$SCRIPT' + fetch_url() { echo \"fetch_url \$*\" >> '$WORK/network.log'; return 1; } + fetch_url_to() { echo \"fetch_url_to \$*\" >> '$WORK/network.log'; return 1; } + main --no-service $* + " +} + +@test "offline install: installs the local binary without any network access" { + run_offline --binary "$WORK/media/maintenant-v1.2.3-linux-amd64" --sha256sums "$WORK/media/SHA256SUMS" --addr 0.0.0.0:8080 + [ "$status" -eq 0 ] + [ ! -e "$WORK/network.log" ] + [ "$(cat "$WORK/bin/maintenant")" = "maintenant binary" ] + grep -qx 'MAINTENANT_ADDR=0.0.0.0:8080' "$WORK/etc/maintenant.env" + [[ "$output" == *"Version : v1.2.3"* ]] +} + +@test "offline install: a renamed binary is found in SHA256SUMS by its checksum" { + cp "$WORK/media/maintenant-v1.2.3-linux-amd64" "$WORK/media/maintenant" + run_offline --binary "$WORK/media/maintenant" --sha256sums "$WORK/media/SHA256SUMS" + [ "$status" -eq 0 ] + [[ "$output" == *"Version : v1.2.3"* ]] +} + +@test "offline install: a checksum mismatch exits 21" { + printf 'tampered' > "$WORK/media/maintenant-v1.2.3-linux-amd64" + run_offline --binary "$WORK/media/maintenant-v1.2.3-linux-amd64" --sha256sums "$WORK/media/SHA256SUMS" + [ "$status" -eq 21 ] + [ ! -e "$WORK/bin/maintenant" ] +} + +@test "offline install: a binary listed for another architecture exits 21" { + (cd "$WORK/media" && sha256sum maintenant-v1.2.3-linux-amd64 \ + | sed 's/linux-amd64/linux-arm64/' > SHA256SUMS) + run_offline --binary "$WORK/media/maintenant-v1.2.3-linux-amd64" --sha256sums "$WORK/media/SHA256SUMS" + [ "$status" -eq 21 ] +} + +@test "offline install: a missing binary exits 2" { + run_offline --binary "$WORK/media/nope" + [ "$status" -eq 2 ] +} + +@test "offline install: without --sha256sums it warns and installs" { + run_offline --binary "$WORK/media/maintenant-v1.2.3-linux-amd64" + [ "$status" -eq 0 ] + [[ "$output" == *"integrity of"*"is not checked"* ]] + [ ! -e "$WORK/network.log" ] +} + +@test "check_prereqs: an offline install needs neither curl nor wget" { + run bash -c " + id() { echo '0'; } + install() { :; } + useradd() { :; } + export -f id install useradd + PATH=/usr/bin/no-such-dir + _INSTALL_SH_TESTING=1 NO_COLOR=1 . '$SCRIPT' + parse_maintenant_flags --no-service --binary ./maintenant + check_prereqs + " + [ "$status" -eq 0 ] +} diff --git a/deploy/install/test/service.bats b/deploy/install/test/service.bats new file mode 100644 index 00000000..8d265749 --- /dev/null +++ b/deploy/install/test/service.bats @@ -0,0 +1,243 @@ +#!/usr/bin/env bats +# Tests for the binary swap, the systemd unit and the directories it may write +bats_require_minimum_version 1.5.0 + +load 'setup' + +SCRIPT="$(cd "$(dirname "$BATS_TEST_FILENAME")/.." && pwd)/install.sh" +UNIT="$(cd "$(dirname "$BATS_TEST_FILENAME")/.." && pwd)/maintenant.service" + +setup() { + WORK=$(mktemp -d) +} + +teardown() { + rm -rf "$WORK" +} + +run_install_service() { + run bash -c " + NO_COLOR=1 + SERVICE_FILE='$WORK/maintenant.service' + export NO_COLOR SERVICE_FILE + systemctl() { + echo \"\$*\" >> '$WORK/systemctl.log' + case \"\$*\" in + start*|restart*|'enable --now'*) touch '$WORK/active' ;; + is-active*) [ -e '$WORK/active' ] ;; + esac + } + sleep() { :; } + journalctl() { :; } + export -f systemctl sleep journalctl + [ -z '${ALREADY_ACTIVE:-}' ] || touch '$WORK/active' + _INSTALL_SH_TESTING=1 . '$SCRIPT' + DATA_DIR=/var/lib/maintenant + DB_DIR=/var/lib/maintenant + install_service + " +} + +@test "install_service: restarts a service that is already running" { + ALREADY_ACTIVE=1 run_install_service + [ "$status" -eq 0 ] + grep -qx 'restart maintenant' "$WORK/systemctl.log" +} + +@test "install_service: starts a service that is not running" { + run_install_service + [ "$status" -eq 0 ] + grep -qx 'start maintenant' "$WORK/systemctl.log" + [ -z "$(grep '^restart' "$WORK/systemctl.log")" ] +} + +@test "install_service: --no-service warns when the running service keeps the old binary" { + NO_SERVICE=1 ALREADY_ACTIVE=1 run_install_service + [ "$status" -eq 0 ] + [[ "$output" == *"still runs the previous binary"* ]] + [ -z "$(grep -E '^(start|restart|enable)' "$WORK/systemctl.log")" ] +} + +@test "install_service: a failing restart exits 31" { + run bash -c " + NO_COLOR=1 + SERVICE_FILE='$WORK/maintenant.service' + export NO_COLOR SERVICE_FILE + systemctl() { + case \"\$*\" in + restart*) return 1 ;; + esac + return 0 + } + export -f systemctl + _INSTALL_SH_TESTING=1 . '$SCRIPT' + DATA_DIR=/var/lib/maintenant + DB_DIR=/var/lib/maintenant + install_service + " + [ "$status" -eq 31 ] +} + +summary_for() { + run bash -c " + MAINTENANT_ADDR=10.9.9.9:1 + export MAINTENANT_ADDR + _INSTALL_SH_TESTING=1 NO_COLOR=1 . '$SCRIPT' + CONFIG_DIR='$WORK/etc' + VERSION=v1.0.0 + parse_maintenant_flags $* + print_summary + " +} + +@test "print_summary: shows the address given as a flag" { + mkdir -p "$WORK/etc" + printf 'MAINTENANT_ADDR=0.0.0.0:7000\n' > "$WORK/etc/maintenant.env" + summary_for --addr 0.0.0.0:9000 + [ "$status" -eq 0 ] + [[ "$output" == *"Listens : http://0.0.0.0:9000"$'\n'* ]] +} + +@test "print_summary: shows the address already in the env file" { + mkdir -p "$WORK/etc" + printf 'MAINTENANT_ADDR="0.0.0.0:7000"\n' > "$WORK/etc/maintenant.env" + summary_for --logLevel debug + [ "$status" -eq 0 ] + [[ "$output" == *"Listens : http://0.0.0.0:7000"$'\n'* ]] +} + +@test "print_summary: falls back to the binary's default address" { + summary_for + [ "$status" -eq 0 ] + [[ "$output" == *"Listens : http://127.0.0.1:8080"$'\n'* ]] +} + +@test "install_binary: writes a temporary file next to the binary, then renames it" { + mkdir -p "$WORK/bin" "$WORK/tmp" + printf 'old' > "$WORK/bin/maintenant" + printf 'new' > "$WORK/tmp/maintenant-v1.0.0-linux-amd64" + + run bash -c " + NO_COLOR=1 + export NO_COLOR + install() { + echo \"\${@: -1}\" >> '$WORK/install.log' + cp \"\${@: -2:1}\" \"\${@: -1}\" + } + export -f install + _INSTALL_SH_TESTING=1 . '$SCRIPT' + INSTALL_DIR='$WORK/bin' + TMPDIR_INSTALL='$WORK/tmp' + VERSION=v1.0.0 + ARCH=amd64 + BINARY_SRC='$WORK/tmp/maintenant-v1.0.0-linux-amd64' + install_binary + " + [ "$status" -eq 0 ] + [ "$(cat "$WORK/bin/maintenant")" = "new" ] + written=$(cat "$WORK/install.log") + [ "$(dirname "$written")" = "$WORK/bin" ] + [ "$written" != "$WORK/bin/maintenant" ] + [ "$(ls -A "$WORK/bin")" = "maintenant" ] +} + +@test "the reference unit is what the script writes with its defaults" { + run bash -c " + unset DATA_DIR MAINTENANT_DATA_DIR CONFIG_DIR MAINTENANT_CONFIG_DIR INSTALL_DIR MAINTENANT_INSTALL_DIR SERVICE_USER + _INSTALL_SH_TESTING=1 . '$SCRIPT' + _env_file_value() { :; } + resolve_paths + render_unit + " + [ "$status" -eq 0 ] + [ "$output" = "$(cat "$UNIT")" ] +} + +render_for() { + run bash -c " + unset DATA_DIR MAINTENANT_DATA_DIR + _INSTALL_SH_TESTING=1 NO_COLOR=1 . '$SCRIPT' + CONFIG_DIR='$WORK/etc' + parse_maintenant_flags $* + resolve_paths + render_unit + " +} + +@test "unit: --db outside the data dir is writable" { + render_for --db /srv/maintenant-db/maintenant.db + [ "$status" -eq 0 ] + [[ "$output" == *"ReadWritePaths=/var/lib/maintenant /srv/maintenant-db"$'\n'* ]] +} + +@test "unit: --data-dir moves the working directory, the writable path and MAINTENANT_DATA_DIR" { + render_for --data-dir /srv/maintenant/ + [ "$status" -eq 0 ] + [[ "$output" == *"WorkingDirectory=/srv/maintenant"$'\n'* ]] + [[ "$output" == *"ReadWritePaths=/srv/maintenant"$'\n'* ]] + [[ "$output" == *"Environment=MAINTENANT_DATA_DIR=/srv/maintenant"$'\n'* ]] +} + +@test "unit: a relative --db lives in the data dir" { + render_for --data-dir /srv/maintenant --db ./db/maintenant.db + [ "$status" -eq 0 ] + [[ "$output" == *"ReadWritePaths=/srv/maintenant /srv/maintenant/db"$'\n'* ]] +} + +@test "unit: paths already in the env file are kept on a rerun without flags" { + mkdir -p "$WORK/etc" + printf 'MAINTENANT_DATA_DIR="/srv/agent"\nMAINTENANT_DB=/srv/db/maintenant.db\n' > "$WORK/etc/maintenant.env" + render_for --logLevel debug + [ "$status" -eq 0 ] + [[ "$output" == *"WorkingDirectory=/srv/agent"$'\n'* ]] + [[ "$output" == *"ReadWritePaths=/srv/agent /srv/db"$'\n'* ]] +} + +@test "resolve_paths: a relative data dir exits 2" { + render_for --data-dir data + [ "$status" -eq 2 ] +} + +@test "resolve_paths: a path the unit file cannot carry exits 2" { + render_for --db '"/srv/my db/maintenant.db"' + [ "$status" -eq 2 ] +} + +@test "prepare_dirs: creates the data and database directories for the service user" { + run bash -c " + NO_COLOR=1 + export NO_COLOR + chown() { echo \"\$*\" >> '$WORK/chown.log'; } + export -f chown + _INSTALL_SH_TESTING=1 . '$SCRIPT' + SERVICE_USER=\$(id -un) + CONFIG_DIR='$WORK/etc' + parse_maintenant_flags --data-dir '$WORK/data' --db '$WORK/db/maintenant.db' + resolve_paths + prepare_dirs + " + [ "$status" -eq 0 ] + [ -d "$WORK/data" ] + [ -d "$WORK/db" ] + grep -q ":.* $WORK/data\$" "$WORK/chown.log" + grep -q ":.* $WORK/db\$" "$WORK/chown.log" +} + +@test "prepare_dirs: refuses a non-empty directory that belongs to someone else" { + mkdir -p "$WORK/shared" + touch "$WORK/shared/other-app.conf" + run bash -c " + NO_COLOR=1 + export NO_COLOR + chown() { echo \"\$*\" >> '$WORK/chown.log'; } + export -f chown + _INSTALL_SH_TESTING=1 . '$SCRIPT' + SERVICE_USER=nobody + CONFIG_DIR='$WORK/etc' + parse_maintenant_flags --data-dir '$WORK/data' --db '$WORK/shared/maintenant.db' + resolve_paths + prepare_dirs + " + [ "$status" -eq 30 ] + [ ! -e "$WORK/chown.log" ] +} diff --git a/deploy/install/test/uninstall.bats b/deploy/install/test/uninstall.bats index 95e7dbcb..b6c04ef1 100644 --- a/deploy/install/test/uninstall.bats +++ b/deploy/install/test/uninstall.bats @@ -99,6 +99,37 @@ run_uninstall() { rm -rf "$FAKE_INSTALL_DIR" "$FAKE_DATA_DIR" "$FAKE_CONFIG_DIR" "$SERVICE_FILE" 2>/dev/null || true } +@test "uninstall --purge: removes the data dir set in the env file" { + FAKE_INSTALL_DIR=$(mktemp -d) + FAKE_DATA_DIR=$(mktemp -d) + FAKE_CONFIG_DIR=$(mktemp -d) + CUSTOM_DATA_DIR=$(mktemp -d) + touch "$CUSTOM_DATA_DIR/agent-identity.json" + printf 'MAINTENANT_DATA_DIR=%s\n' "$CUSTOM_DATA_DIR" > "$FAKE_CONFIG_DIR/maintenant.env" + + run bash -c " + NO_COLOR=1 + id() { echo 0; } + systemctl() { return 0; } + getent() { return 1; } + userdel() { return 0; } + export -f id systemctl getent userdel + INSTALL_DIR='$FAKE_INSTALL_DIR' + DATA_DIR='$FAKE_DATA_DIR' + CONFIG_DIR='$FAKE_CONFIG_DIR' + SERVICE_FILE='/tmp/no-such-service-$$.service' + export INSTALL_DIR DATA_DIR CONFIG_DIR SERVICE_FILE NO_COLOR + _INSTALL_SH_TESTING=1 . '$SCRIPT' + parse_maintenant_flags --uninstall --purge + check_prereqs + uninstall + " + [ "$status" -eq 0 ] + [ ! -d "$CUSTOM_DATA_DIR" ] + + rm -rf "$FAKE_INSTALL_DIR" "$FAKE_DATA_DIR" "$FAKE_CONFIG_DIR" "$CUSTOM_DATA_DIR" +} + @test "uninstall on system without maintenant returns 0 (idempotent)" { FAKE_INSTALL_DIR=$(mktemp -d) FAKE_DATA_DIR=$(mktemp -d) diff --git a/internal/app/app.go b/internal/app/app.go index 0d2c02b2..baa02390 100644 --- a/internal/app/app.go +++ b/internal/app/app.go @@ -725,6 +725,7 @@ func New(cfg Config, logger *slog.Logger, opts ...Option) (*App, error) { // --- Telemetry (SHM SDK, opt-out via MAINTENANT_DISABLE_TELEMETRY) --- a.telemetrySvc = telemetry.New(telemetry.Config{ Disabled: cfg.DisableTelemetry, + DataDir: filepath.Join(filepath.Dir(cfg.DBPath), "shm"), AppVersion: cfg.Version, }, telemetry.Deps{ Containers: containerStore, diff --git a/internal/app/app_test.go b/internal/app/app_test.go index 448b869d..01b01990 100644 --- a/internal/app/app_test.go +++ b/internal/app/app_test.go @@ -1,7 +1,9 @@ package app_test import ( + "bytes" "context" + "encoding/json" "io" "log/slog" "net" @@ -36,6 +38,32 @@ func TestNew_DegradedBoot(t *testing.T) { assert.NotNil(t, a, "app.New() must return a non-nil *App") } +func TestNew_TelemetryDataDirSitsNextToTheDatabase(t *testing.T) { + tmpDir := t.TempDir() + cfg, _ := degradedEnv(t, tmpDir) + var logs bytes.Buffer + logger := slog.New(slog.NewJSONHandler(&logs, nil)) + + _, err := app.New(cfg, logger) + require.NoError(t, err) + + want := filepath.Join(tmpDir, "shm") + assert.DirExists(t, want) + var found bool + for _, line := range bytes.Split(logs.Bytes(), []byte("\n")) { + var rec struct { + Msg string `json:"msg"` + DataDir string `json:"datadir"` + } + if json.Unmarshal(line, &rec) != nil || rec.Msg != "telemetry enabled" { + continue + } + found = true + assert.Equal(t, want, rec.DataDir) + } + assert.True(t, found, "telemetry must start with a directory next to the database") +} + // TestStart_DegradedMode (T008): Start() must launch non-container services and expose // the HTTP server even when the container runtime is unreachable. func TestStart_DegradedMode(t *testing.T) { diff --git a/internal/app/install_script_test.go b/internal/app/install_script_test.go new file mode 100644 index 00000000..76056e4c --- /dev/null +++ b/internal/app/install_script_test.go @@ -0,0 +1,92 @@ +// Copyright 2026 Benjamin Touchard (Kolapsis) +// SPDX-License-Identifier: Apache-2.0 + +package app_test + +import ( + "os" + "os/exec" + "path/filepath" + "regexp" + "strings" + "testing" + + "github.com/stretchr/testify/assert" + "github.com/stretchr/testify/require" + + "github.com/kolapsis/maintenant/internal/app" +) + +var installScriptPath = filepath.Join("..", "..", "deploy", "install", "install.sh") + +func readInstallScript(t *testing.T) string { + t.Helper() + src, err := os.ReadFile(installScriptPath) + require.NoError(t, err) + return string(src) +} + +func scriptWordList(t *testing.T, src, name string) []string { + t.Helper() + start := strings.Index(src, "\n"+name+"=\"") + require.NotEqual(t, -1, start, "%s is not assigned in install.sh", name) + rest := src[start+len(name)+3:] + end := strings.IndexByte(rest, '"') + require.NotEqual(t, -1, end, "%s is not closed in install.sh", name) + return strings.Fields(strings.ReplaceAll(rest[:end], "\\\n", " ")) +} + +func configurationFlags() (bools, values []string) { + for _, spec := range app.Registry { + switch { + case spec.NoEnv: + case spec.Type == app.FlagTypeBool: + bools = append(bools, spec.FlagName) + default: + values = append(values, spec.FlagName) + } + } + return bools, values +} + +func TestInstallScript_AcceptsEveryConfigurationFlag(t *testing.T) { + src := readInstallScript(t) + bools, values := configurationFlags() + + assert.ElementsMatch(t, bools, scriptWordList(t, src, "BOOL_FLAGS"), + "the boolean flags of install.sh must be the binary's FlagTypeBool flags") + assert.ElementsMatch(t, values, scriptWordList(t, src, "VALUE_FLAGS"), + "the value flags of install.sh must be the binary's other flags") +} + +func TestInstallScript_HelpListsEveryConfigurationFlag(t *testing.T) { + src := readInstallScript(t) + start := strings.Index(src, "\nusage() {") + require.NotEqual(t, -1, start) + end := strings.Index(src[start:], "\nEOF\n") + require.NotEqual(t, -1, end) + help := src[start : start+end] + + bools, values := configurationFlags() + for _, name := range append(bools, values...) { + assert.Regexp(t, `(?m)^\s+--`+regexp.QuoteMeta(name)+`(\s|$)`, help, "install.sh --help omits --%s", name) + } +} + +func TestInstallScript_WritesTheEnvNamesTheBinaryReads(t *testing.T) { + var names, want []string + for _, spec := range app.Registry { + if spec.NoEnv { + continue + } + names = append(names, spec.FlagName) + want = append(want, spec.EnvName) + } + + args := append([]string{"-c", `. "$0"; for f in "$@"; do flag_to_env "$f"; echo; done`, installScriptPath}, names...) + cmd := exec.Command("sh", args...) + cmd.Env = append(os.Environ(), "_INSTALL_SH_TESTING=1", "NO_COLOR=1") + out, err := cmd.Output() + require.NoError(t, err) + assert.Equal(t, want, strings.Fields(string(out))) +} diff --git a/internal/app/main_test.go b/internal/app/main_test.go index 265b0d68..ff12cc61 100644 --- a/internal/app/main_test.go +++ b/internal/app/main_test.go @@ -13,5 +13,9 @@ import ( func TestMain(m *testing.M) { extension.Register(tiers.Policy{}, nil) + // Tests that call App.Start must never report to the real telemetry endpoint. + if err := os.Setenv("DO_NOT_TRACK", "1"); err != nil { + panic(err) + } os.Exit(m.Run()) } diff --git a/internal/telemetry/telemetry.go b/internal/telemetry/telemetry.go index d8596409..5844b898 100644 --- a/internal/telemetry/telemetry.go +++ b/internal/telemetry/telemetry.go @@ -16,7 +16,6 @@ import ( // operator-configurable; the cfg struct exists so tests can override them. const ( defaultEndpoint = "https://metrics.kolapsis.com" - defaultDataDir = "/data/shm" defaultEnvironment = "production" appName = "Maintenant" ) @@ -28,11 +27,10 @@ const ( ) // Config carries the wire-relevant constants for the telemetry subsystem. -// All fields have safe defaults populated by New when unset. type Config struct { Disabled bool // true if MAINTENANT_DISABLE_TELEMETRY is truthy Endpoint string // default: defaultEndpoint - DataDir string // default: defaultDataDir + DataDir string // no default: telemetry stays off when empty or unwritable AppVersion string // injected via ldflags at build time Environment string // default: defaultEnvironment } @@ -152,9 +150,6 @@ func applyDefaults(cfg Config) Config { if cfg.Endpoint == "" { cfg.Endpoint = defaultEndpoint } - if cfg.DataDir == "" { - cfg.DataDir = defaultDataDir - } if cfg.Environment == "" { cfg.Environment = defaultEnvironment } From e2bc989d676a2c32d77af6ee8213a8935f67761a Mon Sep 17 00:00:00 2001 From: Benjamin Date: Wed, 30 Sep 2026 16:24:47 +0200 Subject: [PATCH 02/54] fix(updates): check official images, real rollbacks, apply update labels - Check Docker Hub official images on latest; skip only images the local Docker runtime reports as never pulled (no RepoDigests, read from the image list at scan time). - Rollback commands return to the image the container ran before the update: running repo digest, else digest baseline, else fixed version tag. Compose retags the unchanged tag or pins the previous image, Kubernetes sets the image explicitly. previous_digest is persisted. - The Compose update command asks to write the new tag in the compose file when the tag changes. - Apply maintenant.update.track, ignore_major, digest_only and alert_on. - Image exclusions match the image name without tag or digest. - Tag filters keep the current tag as the digest comparison reference. --- internal/api/v1/updates.go | 4 +- internal/api/v1/updates_test.go | 63 ++++++ internal/app/app.go | 16 +- internal/app/wiring.go | 19 +- internal/app/wiring_update_test.go | 25 +++ internal/docker/discovery.go | 25 +++ internal/docker/repodigests_test.go | 51 +++++ internal/docker/runtime.go | 5 + internal/update/container_adapter.go | 65 +++--- internal/update/container_adapter_test.go | 57 ++++++ internal/update/labels.go | 42 ++-- internal/update/model.go | 31 +++ internal/update/scanner.go | 230 +++++++++++++++------- internal/update/scanner_rules_test.go | 196 ++++++++++++++++++ internal/update/scanner_test.go | 19 +- internal/update/semver.go | 35 +++- internal/update/service.go | 113 +++++++---- internal/update/service_commands_test.go | 205 +++++++++++++++++++ 18 files changed, 1027 insertions(+), 174 deletions(-) create mode 100644 internal/docker/repodigests_test.go create mode 100644 internal/update/container_adapter_test.go create mode 100644 internal/update/scanner_rules_test.go create mode 100644 internal/update/service_commands_test.go diff --git a/internal/api/v1/updates.go b/internal/api/v1/updates.go index 1d650c3a..f4f3e30f 100644 --- a/internal/api/v1/updates.go +++ b/internal/api/v1/updates.go @@ -175,8 +175,8 @@ func (h *UpdateHandler) HandleGetContainerUpdate(w http.ResponseWriter, r *http. ci, ciErr := h.containers.GetContainerInfo(r.Context(), containerID) if ciErr == nil { resp["update_command"] = h.service.GenerateUpdateCommand(ci, u.LatestTag) - if u.PreviousDigest != "" { - resp["rollback_command"] = h.service.GenerateRollbackCommand(ci, u.PreviousDigest) + if cmd := h.service.GenerateRollbackCommand(ci, u); cmd != "" { + resp["rollback_command"] = cmd } // Tag filter labels (raw pattern strings, shown as configured even if regex was invalid) if v := ci.Labels["maintenant.update.tag-include"]; v != "" { diff --git a/internal/api/v1/updates_test.go b/internal/api/v1/updates_test.go index 230bf491..736fc734 100644 --- a/internal/api/v1/updates_test.go +++ b/internal/api/v1/updates_test.go @@ -169,3 +169,66 @@ func TestHandleListAgents_CarriesOSSupport(t *testing.T) { } t.Fatal("agent web-03 missing from the listing") } + +type fixedContainerInfo struct { + info update.ContainerInfo +} + +func (f fixedContainerInfo) GetContainerInfo(context.Context, string) (update.ContainerInfo, error) { + return f.info, nil +} + +func TestHandleGetContainerUpdate_RollbackReturnsToThePreviousImage(t *testing.T) { + ctx := context.Background() + logger := slog.New(slog.NewTextHandler(io.Discard, nil)) + db := storetest.Open(t, logger) + updateStore := store.NewUpdateStore(db) + + scanID, err := updateStore.InsertScanRecord(ctx, &update.ScanRecord{StartedAt: time.Now(), Status: update.ScanStatusCompleted}) + require.NoError(t, err) + _, err = updateStore.InsertImageUpdate(ctx, &update.ImageUpdate{ + ScanID: scanID, + ContainerID: "ctr1", + ContainerName: "app-web-1", + Image: "nginx:latest", + CurrentTag: "latest", + CurrentDigest: "sha256:old", + Registry: "registry-1.docker.io", + LatestTag: "latest", + LatestDigest: "sha256:new", + UpdateType: update.UpdateTypeDigestOnly, + PreviousDigest: "sha256:old", + Status: update.StatusAvailable, + DetectedAt: time.Now(), + }) + require.NoError(t, err) + + h := NewUpdateHandler(update.NewService(update.Deps{ + Store: updateStore, + Scanner: update.NewScanner(update.NewRegistryClient(), updateStore, logger), + Containers: noContainers{}, + Logger: logger, + }), updateStore, fixedContainerInfo{info: update.ContainerInfo{ + ExternalID: "ctr1", + Name: "app-web-1", + Image: "nginx:latest", + OrchestrationGroup: "app", + OrchestrationUnit: "web", + ComposeWorkingDir: "/srv/app", + RuntimeType: "docker", + }}) + + req := httptest.NewRequest(http.MethodGet, "/api/v1/updates/ctr1", nil) + req.SetPathValue("container_id", "ctr1") + rec := httptest.NewRecorder() + h.HandleGetContainerUpdate(rec, req) + require.Equal(t, http.StatusOK, rec.Code) + + var body struct { + UpdateCommand string `json:"update_command"` + RollbackCommand string `json:"rollback_command"` + } + require.NoError(t, json.Unmarshal(rec.Body.Bytes(), &body)) + assert.NotEqual(t, body.UpdateCommand, body.RollbackCommand) + assert.Contains(t, body.RollbackCommand, "docker tag nginx@sha256:old nginx:latest") +} diff --git a/internal/app/app.go b/internal/app/app.go index baa02390..2ea365c8 100644 --- a/internal/app/app.go +++ b/internal/app/app.go @@ -528,10 +528,11 @@ func New(cfg Config, logger *slog.Logger, opts ...Option) (*App, error) { registryClient := update.NewRegistryClient() updateScanner := update.NewScanner(registryClient, updateStore, logger) containerAdapter := update.NewContainerServiceAdapter(a.containerSvc) - // Wire live label fetching from Docker so maintenant.update.tag-include/exclude labels - // are available at scan time (labels are not persisted in SQLite). + // Wire live label and image digest fetching from Docker so maintenant.update.* labels + // and the running image's digests are available at scan time (neither is persisted). if dr, ok := a.rt.(*docker.Runtime); ok { - containerAdapter.WithLabelFetcher(&dockerLabelFetcher{rt: dr}) + fetcher := &dockerLabelFetcher{rt: dr} + containerAdapter.WithLabelFetcher(fetcher).WithRepoDigestFetcher(fetcher) } var updateEnricher update.Enricher @@ -1103,13 +1104,16 @@ func (a *App) swarmNodeStoreAsInterface() swarm.NodeStore { return a.swarmNodeStore } -// dockerLabelFetcher implements update.LabelFetcher for Docker runtimes. -// It fetches live container labels at scan time so tag-include/tag-exclude labels -// are available without persisting them in SQLite. +// dockerLabelFetcher implements update.LabelFetcher and update.RepoDigestFetcher for Docker runtimes. +// It fetches live container labels and image digests at scan time, without persisting them. type dockerLabelFetcher struct { rt *docker.Runtime } +func (f *dockerLabelFetcher) FetchRepoDigests(ctx context.Context) (map[string][]string, error) { + return f.rt.ContainerRepoDigests(ctx) +} + func (f *dockerLabelFetcher) FetchLabels(ctx context.Context) (map[string]map[string]string, error) { results, err := f.rt.DiscoverAllWithLabels(ctx) if err != nil { diff --git a/internal/app/wiring.go b/internal/app/wiring.go index 8fa9b7a3..c32ef472 100644 --- a/internal/app/wiring.go +++ b/internal/app/wiring.go @@ -16,6 +16,7 @@ import ( "github.com/kolapsis/maintenant/internal/extension" "github.com/kolapsis/maintenant/internal/heartbeat" "github.com/kolapsis/maintenant/internal/security" + "github.com/kolapsis/maintenant/internal/update" ) const restartLoopAlertType = "restart_loop" @@ -484,6 +485,19 @@ func updateDetectedAlert(m map[string]any, withChangelog bool) alert.Event { } } +// updateAlertWanted applies the container's maintenant.update.alert_on label to the severity computed for its update. +func updateAlertWanted(m map[string]any, severity string) bool { + alertOn, _ := m["alert_on"].(string) + switch alertOn { + case update.AlertOnNone: + return false + case update.AlertOnCritical: + return severity == alert.SeverityCritical + default: + return true + } +} + // updateResolvedAlert builds the recovery event when a container's update is no // longer pending. EntityID uses the same container UID as updateDetectedAlert so // the right alert is resolved by dedup key. The caller sets Timestamp. @@ -532,7 +546,10 @@ func (a *App) wireUpdateCallback() { } sendAlert(updateResolvedAlert(m)) case event.UpdateDetected: - sendAlert(updateDetectedAlert(m, extension.Allows(extension.CapChangelog))) + evt := updateDetectedAlert(m, extension.Allows(extension.CapChangelog)) + if updateAlertWanted(m, evt.Severity) { + sendAlert(evt) + } } }) } diff --git a/internal/app/wiring_update_test.go b/internal/app/wiring_update_test.go index 7d28ef99..d870bce1 100644 --- a/internal/app/wiring_update_test.go +++ b/internal/app/wiring_update_test.go @@ -88,6 +88,31 @@ func TestUpdateDetectedAlert_SeverityFromRiskScore(t *testing.T) { } } +func TestUpdateAlertWanted_FollowsAlertOnLabel(t *testing.T) { + cases := []struct { + alertOn string + risk int + want bool + }{ + {"", 10, true}, + {"all", 10, true}, + {"critical", 90, true}, + {"critical", 70, false}, + {"critical", 10, false}, + {"none", 90, false}, + } + for _, tc := range cases { + m := detectedPayload("svc", "id", tc.risk) + if tc.alertOn != "" { + m["alert_on"] = tc.alertOn + } + evt := updateDetectedAlert(m, false) + if got := updateAlertWanted(m, evt.Severity); got != tc.want { + t.Errorf("alert_on=%q risk=%d → %v, want %v", tc.alertOn, tc.risk, got, tc.want) + } + } +} + func TestUpdateDetectedAlert_ProDetailsGated(t *testing.T) { ce := updateDetectedAlert(detectedPayload("svc", "id", 90), false) if _, ok := ce.Details["update_command"]; ok { diff --git a/internal/docker/discovery.go b/internal/docker/discovery.go index 4aa74edc..615e00c8 100644 --- a/internal/docker/discovery.go +++ b/internal/docker/discovery.go @@ -128,6 +128,31 @@ func (c *Client) DiscoverAllWithLabels(ctx context.Context) ([]*DiscoveryResult, return results, nil } +// ContainerRepoDigests maps each container ID to the repo digests of the image it runs, empty for an image never pulled nor pushed. +func (c *Client) ContainerRepoDigests(ctx context.Context) (map[string][]string, error) { + containers, err := c.cli.ContainerList(ctx, client.ContainerListOptions{All: true}) + if err != nil { + return nil, fmt.Errorf("container list: %w", err) + } + images, err := c.cli.ImageList(ctx, client.ImageListOptions{}) + if err != nil { + return nil, fmt.Errorf("image list: %w", err) + } + + digestsByImage := make(map[string][]string, len(images.Items)) + for _, img := range images.Items { + digestsByImage[img.ID] = img.RepoDigests + } + + out := make(map[string][]string, len(containers.Items)) + for _, dc := range containers.Items { + if digests, ok := digestsByImage[dc.ImageID]; ok { + out[dc.ID] = digests + } + } + return out, nil +} + // inspectResult holds the mapped container along with its extracted security config. type inspectResult struct { Container *cmodel.Container diff --git a/internal/docker/repodigests_test.go b/internal/docker/repodigests_test.go new file mode 100644 index 00000000..c09351eb --- /dev/null +++ b/internal/docker/repodigests_test.go @@ -0,0 +1,51 @@ +// Copyright 2026 Benjamin Touchard (Kolapsis) +// SPDX-License-Identifier: Apache-2.0 + +package docker + +import ( + "context" + "io" + "log/slog" + "testing" + + "github.com/moby/moby/api/types/container" + "github.com/moby/moby/api/types/image" + "github.com/moby/moby/client" + "github.com/stretchr/testify/assert" + "github.com/stretchr/testify/require" +) + +type imageFakeAPI struct { + *fakeAPI + images []image.Summary +} + +func (f *imageFakeAPI) ImageList(context.Context, client.ImageListOptions) (client.ImageListResult, error) { + return client.ImageListResult{Items: f.images}, nil +} + +func TestContainerRepoDigests(t *testing.T) { + api := &imageFakeAPI{ + fakeAPI: &fakeAPI{list: []container.Summary{ + {ID: "pulled", ImageID: "sha256:img-pulled"}, + {ID: "built", ImageID: "sha256:img-built"}, + {ID: "gone", ImageID: "sha256:img-gone"}, + }}, + images: []image.Summary{ + {ID: "sha256:img-pulled", RepoDigests: []string{"nginx@sha256:abc"}}, + {ID: "sha256:img-built", RepoTags: []string{"myapp:latest"}}, + }, + } + c := &Client{cli: api, logger: slog.New(slog.NewTextHandler(io.Discard, nil))} + + got, err := c.ContainerRepoDigests(context.Background()) + require.NoError(t, err) + + assert.Equal(t, []string{"nginx@sha256:abc"}, got["pulled"]) + built, ok := got["built"] + assert.True(t, ok) + assert.Empty(t, built) + _, ok = got["gone"] + assert.False(t, ok, "a container whose image is not listed stays unknown") +} diff --git a/internal/docker/runtime.go b/internal/docker/runtime.go index ecab517c..b9c927ff 100644 --- a/internal/docker/runtime.go +++ b/internal/docker/runtime.go @@ -90,6 +90,11 @@ func (r *Runtime) DiscoverAllWithLabels(ctx context.Context) ([]*DiscoveryResult return r.client.DiscoverAllWithLabels(ctx) } +// ContainerRepoDigests delegates to the underlying Docker client. +func (r *Runtime) ContainerRepoDigests(ctx context.Context) (map[string][]string, error) { + return r.client.ContainerRepoDigests(ctx) +} + func (r *Runtime) StreamEvents(ctx context.Context) <-chan runtime.RuntimeEvent { dockerCh := r.client.StreamEvents(ctx) out := make(chan runtime.RuntimeEvent, 64) diff --git a/internal/update/container_adapter.go b/internal/update/container_adapter.go index 18facf16..7472bb8d 100644 --- a/internal/update/container_adapter.go +++ b/internal/update/container_adapter.go @@ -17,10 +17,16 @@ type LabelFetcher interface { FetchLabels(ctx context.Context) (map[string]map[string]string, error) } +// RepoDigestFetcher maps each container external ID to the repo digests of its image, an empty list marking an image that never came from a registry. +type RepoDigestFetcher interface { + FetchRepoDigests(ctx context.Context) (map[string][]string, error) +} + // ContainerServiceAdapter adapts container.Service to the ContainerLister interface. type ContainerServiceAdapter struct { - svc *container.Service - labelFetcher LabelFetcher // optional — nil when runtime doesn't support label fetching + svc *container.Service + labelFetcher LabelFetcher // optional, nil when runtime doesn't support label fetching + digestFetcher RepoDigestFetcher // optional, nil when the runtime does not report image digests } // NewContainerServiceAdapter creates a new adapter. @@ -35,6 +41,12 @@ func (a *ContainerServiceAdapter) WithLabelFetcher(lf LabelFetcher) *ContainerSe return a } +// WithRepoDigestFetcher attaches the runtime source of ContainerInfo.RepoDigests and ContainerInfo.LocallyBuilt. +func (a *ContainerServiceAdapter) WithRepoDigestFetcher(df RepoDigestFetcher) *ContainerServiceAdapter { + a.digestFetcher = df + return a +} + // ListContainerInfos returns container info for all running containers. func (a *ContainerServiceAdapter) ListContainerInfos(ctx context.Context) ([]ContainerInfo, error) { containers, err := a.svc.ListContainers(ctx, container.ListContainersOpts{ @@ -49,24 +61,21 @@ func (a *ContainerServiceAdapter) ListContainerInfos(ctx context.Context) ([]Con if a.labelFetcher != nil { labelsByExtID, _ = a.labelFetcher.FetchLabels(ctx) } + var digestsByExtID map[string][]string + if a.digestFetcher != nil { + digestsByExtID, _ = a.digestFetcher.FetchRepoDigests(ctx) + } infos := make([]ContainerInfo, 0, len(containers)) for _, c := range containers { if c.IsIgnored || c.Archived { continue } - infos = append(infos, ContainerInfo{ - UID: c.ID, - ExternalID: c.ExternalID, - Name: c.Name, - Image: c.Image, - Labels: labelsByExtID[c.ExternalID], - OrchestrationGroup: c.OrchestrationGroup, - OrchestrationUnit: c.OrchestrationUnit, - RuntimeType: c.RuntimeType, - ControllerKind: c.ControllerKind, - ComposeWorkingDir: c.ComposeWorkingDir, - }) + info := newContainerInfo(c, labelsByExtID[c.ExternalID]) + digests, known := digestsByExtID[c.ExternalID] + info.RepoDigests = digests + info.LocallyBuilt = known && len(digests) == 0 + infos = append(infos, info) } return infos, nil } @@ -86,19 +95,23 @@ func (a *ContainerServiceAdapter) GetContainerInfo(ctx context.Context, external for _, c := range containers { if c.ExternalID == externalID { - return ContainerInfo{ - UID: c.ID, - ExternalID: c.ExternalID, - Name: c.Name, - Image: c.Image, - Labels: labelsByExtID[c.ExternalID], - OrchestrationGroup: c.OrchestrationGroup, - OrchestrationUnit: c.OrchestrationUnit, - RuntimeType: c.RuntimeType, - ControllerKind: c.ControllerKind, - ComposeWorkingDir: c.ComposeWorkingDir, - }, nil + return newContainerInfo(c, labelsByExtID[c.ExternalID]), nil } } return ContainerInfo{}, fmt.Errorf("container not found: %s", externalID) } + +func newContainerInfo(c *container.Container, labels map[string]string) ContainerInfo { + return ContainerInfo{ + UID: c.ID, + ExternalID: c.ExternalID, + Name: c.Name, + Image: c.Image, + Labels: labels, + OrchestrationGroup: c.OrchestrationGroup, + OrchestrationUnit: c.OrchestrationUnit, + RuntimeType: c.RuntimeType, + ControllerKind: c.ControllerKind, + ComposeWorkingDir: c.ComposeWorkingDir, + } +} diff --git a/internal/update/container_adapter_test.go b/internal/update/container_adapter_test.go new file mode 100644 index 00000000..93a0cacb --- /dev/null +++ b/internal/update/container_adapter_test.go @@ -0,0 +1,57 @@ +// Copyright 2026 Benjamin Touchard (kOlapsis) +// SPDX-License-Identifier: Apache-2.0 + +package update + +import ( + "context" + "testing" + + "github.com/stretchr/testify/assert" + "github.com/stretchr/testify/require" + + "github.com/kolapsis/maintenant/internal/container" +) + +type fakeContainerStore struct { + container.ContainerStore + list []*container.Container +} + +func (f fakeContainerStore) ListContainers(context.Context, container.ListContainersOpts) ([]*container.Container, error) { + return f.list, nil +} + +type fakeRepoDigests map[string][]string + +func (f fakeRepoDigests) FetchRepoDigests(context.Context) (map[string][]string, error) { + return f, nil +} + +func TestContainerServiceAdapter_RepoDigestsAndLocalBuilds(t *testing.T) { + svc := container.NewService(container.Deps{ + Store: fakeContainerStore{list: []*container.Container{ + {ExternalID: "pulled", Name: "web", Image: "nginx:latest"}, + {ExternalID: "built", Name: "app", Image: "myapp:latest"}, + {ExternalID: "unknown", Name: "remote", Image: "redis:latest"}, + }}, + Logger: testLogger(), + }) + adapter := NewContainerServiceAdapter(svc).WithRepoDigestFetcher(fakeRepoDigests{ + "pulled": {"nginx@sha256:abc"}, + "built": {}, + }) + + infos, err := adapter.ListContainerInfos(context.Background()) + require.NoError(t, err) + require.Len(t, infos, 3) + + byID := map[string]ContainerInfo{} + for _, info := range infos { + byID[info.ExternalID] = info + } + assert.Equal(t, []string{"nginx@sha256:abc"}, byID["pulled"].RepoDigests) + assert.False(t, byID["pulled"].LocallyBuilt) + assert.True(t, byID["built"].LocallyBuilt) + assert.False(t, byID["unknown"].LocallyBuilt, "a runtime that says nothing about the image must not hide it") +} diff --git a/internal/update/labels.go b/internal/update/labels.go index 5bb19ba6..a1205bfd 100644 --- a/internal/update/labels.go +++ b/internal/update/labels.go @@ -15,7 +15,7 @@ const updateLabelPrefix = "maintenant.update." func ParseUpdateLabels(labels map[string]string, logger *slog.Logger) UpdateConfig { cfg := UpdateConfig{ Enabled: true, // enabled by default - AlertOn: "all", + AlertOn: AlertOnAll, } for key, value := range labels { @@ -27,17 +27,14 @@ func ParseUpdateLabels(labels map[string]string, logger *slog.Logger) UpdateConf switch suffix { case "enabled": - switch strings.ToLower(value) { - case "false", "0", "no": - cfg.Enabled = false - case "true", "1", "yes": - cfg.Enabled = true - default: + if b, ok := parseLabelBool(value); ok { + cfg.Enabled = b + } else { logger.Warn("invalid maintenant.update.enabled value", "value", value) } case "track": switch strings.ToLower(value) { - case "major", "minor", "patch", "digest": + case TrackMajor, TrackMinor, TrackPatch, TrackDigest: cfg.Track = strings.ToLower(value) default: logger.Warn("invalid maintenant.update.track value", "value", value) @@ -45,25 +42,25 @@ func ParseUpdateLabels(labels map[string]string, logger *slog.Logger) UpdateConf case "pin": cfg.Pin = value case "ignore_major": - switch strings.ToLower(value) { - case "true", "1", "yes": - cfg.IgnoreMajor = true - case "false", "0", "no": - cfg.IgnoreMajor = false + if b, ok := parseLabelBool(value); ok { + cfg.IgnoreMajor = b + } else { + logger.Warn("invalid maintenant.update.ignore_major value", "value", value) } case "registry": cfg.Registry = value case "alert_on": switch strings.ToLower(value) { - case "all", "critical", "none": + case AlertOnAll, AlertOnCritical, AlertOnNone: cfg.AlertOn = strings.ToLower(value) default: logger.Warn("invalid maintenant.update.alert_on value", "value", value) } case "digest_only": - switch strings.ToLower(value) { - case "true", "1", "yes": - cfg.DigestOnly = true + if b, ok := parseLabelBool(value); ok { + cfg.DigestOnly = b + } else { + logger.Warn("invalid maintenant.update.digest_only value", "value", value) } case "tag-include": if value == "" { @@ -92,3 +89,14 @@ func ParseUpdateLabels(labels map[string]string, logger *slog.Logger) UpdateConf return cfg } + +func parseLabelBool(value string) (bool, bool) { + switch strings.ToLower(value) { + case "true", "1", "yes": + return true, true + case "false", "0", "no": + return false, true + default: + return false, false + } +} diff --git a/internal/update/model.go b/internal/update/model.go index 41b66bd1..4c6377e1 100644 --- a/internal/update/model.go +++ b/internal/update/model.go @@ -142,6 +142,7 @@ type UpdateResult struct { HasBreakingChanges bool SourceURL string PreviousDigest string + AlertOn string } // ScanError represents an error scanning a specific container. @@ -306,6 +307,21 @@ type EcosystemResult struct { DetectionMethod string `json:"detection_method"` } +// Accepted values of the maintenant.update.track label. +const ( + TrackMajor = "major" + TrackMinor = "minor" + TrackPatch = "patch" + TrackDigest = "digest" +) + +// Accepted values of the maintenant.update.alert_on label. +const ( + AlertOnAll = "all" + AlertOnCritical = "critical" + AlertOnNone = "none" +) + // UpdateConfig holds parsed maintenant.update.* label values. type UpdateConfig struct { Enabled bool @@ -318,3 +334,18 @@ type UpdateConfig struct { TagInclude *regexp.Regexp // compiled tag-include regex, nil if absent/invalid TagExclude *regexp.Regexp // compiled tag-exclude regex, nil if absent/invalid } + +// TrackLevel returns the widest change the labels allow: a Track* value, digest_only and ignore_major applied. +func (cfg UpdateConfig) TrackLevel() string { + if cfg.DigestOnly { + return TrackDigest + } + level := cfg.Track + if level == "" { + level = TrackMajor + } + if cfg.IgnoreMajor && level == TrackMajor { + return TrackMinor + } + return level +} diff --git a/internal/update/scanner.go b/internal/update/scanner.go index fb3a40cb..8cd68d6f 100644 --- a/internal/update/scanner.go +++ b/internal/update/scanner.go @@ -8,6 +8,7 @@ import ( "fmt" "log/slog" "path/filepath" + "slices" "strings" "time" ) @@ -24,6 +25,8 @@ type ContainerInfo struct { RuntimeType string ControllerKind string ComposeWorkingDir string + RepoDigests []string // "repo@sha256:..." of the running image, as the local runtime reports them + LocallyBuilt bool // the runtime reports an image never pulled from nor pushed to a registry } // Scanner checks containers for available updates by comparing tags and digests. @@ -118,10 +121,8 @@ func (sc *Scanner) scanContainer(ctx context.Context, c ContainerInfo, exclusion return nil, fmt.Errorf("cannot parse image reference: %s", c.Image) } - // Skip local/private images that have no registry and no slash (locally built) - if !strings.Contains(imageRef, "/") && currentTag == "latest" && registry == "registry-1.docker.io" { - // Likely a locally-built image (e.g. "myapp" or "myapp:latest") — skip silently - sc.logger.Debug("scanner: skipping likely local image", "image", c.Image) + if c.LocallyBuilt { + sc.logger.Debug("scanner: skipping locally built image", "image", c.Image) return nil, nil } @@ -162,11 +163,24 @@ func (sc *Scanner) scanContainer(ctx context.Context, c ContainerInfo, exclusion fullRef = "library/" + imageRef } + target := scanTarget{ + container: c, + fullRef: fullRef, + currentTag: currentTag, + registry: registry, + runningDigest: repoDigestFor(c.RepoDigests, imageRef), + alertOn: cfg.AlertOn, + } + + level := cfg.TrackLevel() + if level == TrackDigest { + return sc.checkDigest(ctx, target) + } + // List all tags from registry tags, err := sc.registry.ListTags(ctx, fullRef) if err != nil { - // Skip images that fail auth (private/local images not on any registry) - if strings.Contains(err.Error(), "UNAUTHORIZED") || strings.Contains(err.Error(), "NAME_UNKNOWN") || strings.Contains(err.Error(), "denied") { + if isUnreachableImage(err) { sc.logger.Debug("scanner: skipping unreachable image", "image", c.Image, "reason", err.Error()) return nil, nil } @@ -181,70 +195,31 @@ func (sc *Scanner) scanContainer(ctx context.Context, c ContainerInfo, exclusion if _, parseErr := ParseTag(versionPart); parseErr == nil { // Semver tag: apply user-configured tag filters tf := NewTagFilter(cfg.TagInclude, cfg.TagExclude, variant) - tags = tf.Filter(tags) - if len(tags) == 0 { - sc.logger.Warn("scanner: tag filter produced no candidates, skipping update check", + filtered := tf.Filter(tags) + if len(filtered) == 0 { + sc.logger.Warn("scanner: tag filter produced no candidates", "container", c.Name, "image", c.Image) + } + // The current tag stays the reference for digest comparison even when the filter rejects it. + if slices.Contains(tags, currentTag) && !slices.Contains(filtered, currentTag) { + filtered = append(filtered, currentTag) + } + tags = filtered + if len(tags) == 0 { return nil, nil } } // Find best update - bestTag, updateType := FindBestUpdate(currentTag, tags) + bestTag, updateType := findBestUpdate(currentTag, tags, level) if bestTag == "" { return nil, nil } - // Digest-only mode: floating tags, i.e. non-semver channels like "lts", "alpine", - // "stable", "latest", plus partial versions like "v3" or "1.2" that the registry moves. - // Compare the current remote digest against the stored baseline to detect rebuilds. + // Floating tags, i.e. non-semver channels like "lts", "alpine", "stable", "latest", + // plus partial versions like "v3" or "1.2" that the registry moves. if bestTag == currentTag && updateType == UpdateTypeDigestOnly { - tagRef := fullRef + ":" + currentTag - remoteDigest, err := sc.registry.GetDigest(ctx, tagRef) - if err != nil || remoteDigest == "" { - sc.logger.Debug("scanner: cannot fetch digest for channel tag", - "container", c.Name, "tag", currentTag, "error", err) - return nil, nil - } - - baseline, _ := sc.store.GetDigestBaseline(ctx, c.ExternalID) - - // Store/update the baseline for next scan comparison - now := time.Now() - if err := sc.store.UpsertDigestBaseline(ctx, &DigestBaseline{ - ContainerID: c.ExternalID, - Image: c.Image, - Tag: currentTag, - RemoteDigest: remoteDigest, - CheckedAt: now, - }); err != nil { - sc.logger.Warn("scanner: failed to store digest baseline", - "container", c.Name, "error", err) - } - - if baseline == nil || baseline.RemoteDigest == remoteDigest { - // First scan or digest unchanged — no update - return nil, nil - } - - // Digest changed — the tag was republished with a new build - sc.logger.Info("scanner: digest change detected for channel tag", - "container", c.Name, "tag", currentTag, - "old_digest", baseline.RemoteDigest[:19], "new_digest", remoteDigest[:19]) - - return &UpdateResult{ - ContainerID: c.ExternalID, - ContainerName: c.Name, - Image: c.Image, - CurrentTag: currentTag, - CurrentDigest: baseline.RemoteDigest, - PreviousDigest: baseline.RemoteDigest, - Registry: registry, - LatestTag: currentTag, - LatestDigest: remoteDigest, - UpdateType: UpdateTypeDigestOnly, - HasUpdate: true, - }, nil + return sc.checkDigest(ctx, target) } // Semver update: a newer version tag exists @@ -256,26 +231,141 @@ func (sc *Scanner) scanContainer(ctx context.Context, c ContainerInfo, exclusion } result := &UpdateResult{ - ContainerID: c.ExternalID, - ContainerName: c.Name, - Image: c.Image, - CurrentTag: currentTag, - Registry: registry, - LatestTag: bestTag, - LatestDigest: latestDigest, - UpdateType: updateType, - HasUpdate: true, + ContainerID: c.ExternalID, + ContainerName: c.Name, + Image: c.Image, + CurrentTag: currentTag, + Registry: registry, + LatestTag: bestTag, + LatestDigest: latestDigest, + UpdateType: updateType, + HasUpdate: true, + PreviousDigest: target.runningDigest, + AlertOn: cfg.AlertOn, } return result, nil } +// scanTarget carries what the digest comparison needs about one container. +type scanTarget struct { + container ContainerInfo + fullRef string + currentTag string + registry string + runningDigest string + alertOn string +} + +// checkDigest compares the remote digest of the current tag with the one recorded at the +// previous scan, which detects a republished tag without looking at other tags. +func (sc *Scanner) checkDigest(ctx context.Context, t scanTarget) (*UpdateResult, error) { + c := t.container + remoteDigest, err := sc.registry.GetDigest(ctx, t.fullRef+":"+t.currentTag) + if err != nil { + if isUnreachableImage(err) { + sc.logger.Debug("scanner: skipping unreachable tag", + "container", c.Name, "tag", t.currentTag, "reason", err.Error()) + return nil, nil + } + return nil, fmt.Errorf("get digest: %w", err) + } + if remoteDigest == "" { + return nil, nil + } + + baseline, err := sc.store.GetDigestBaseline(ctx, c.ExternalID) + if err != nil { + return nil, fmt.Errorf("get digest baseline: %w", err) + } + + if err := sc.store.UpsertDigestBaseline(ctx, &DigestBaseline{ + ContainerID: c.ExternalID, + Image: c.Image, + Tag: t.currentTag, + RemoteDigest: remoteDigest, + CheckedAt: time.Now(), + }); err != nil { + sc.logger.Warn("scanner: failed to store digest baseline", + "container", c.Name, "error", err) + } + + if baseline == nil || baseline.RemoteDigest == remoteDigest { + return nil, nil + } + + sc.logger.Info("scanner: digest change detected for channel tag", + "container", c.Name, "tag", t.currentTag, + "old_digest", shortDigest(baseline.RemoteDigest), "new_digest", shortDigest(remoteDigest)) + + previous := t.runningDigest + if previous == "" { + previous = baseline.RemoteDigest + } + + return &UpdateResult{ + ContainerID: c.ExternalID, + ContainerName: c.Name, + Image: c.Image, + CurrentTag: t.currentTag, + CurrentDigest: baseline.RemoteDigest, + PreviousDigest: previous, + Registry: t.registry, + LatestTag: t.currentTag, + LatestDigest: remoteDigest, + UpdateType: UpdateTypeDigestOnly, + HasUpdate: true, + AlertOn: t.alertOn, + }, nil +} + +// isUnreachableImage reports a registry answer meaning the image or tag cannot be +// checked there at all (private, unknown or never published), as opposed to a failure. +func isUnreachableImage(err error) bool { + msg := err.Error() + return strings.Contains(msg, "UNAUTHORIZED") || strings.Contains(msg, "NAME_UNKNOWN") || + strings.Contains(msg, "MANIFEST_UNKNOWN") || strings.Contains(msg, "denied") +} + +func shortDigest(d string) string { + if len(d) > 19 { + return d[:19] + } + return d +} + +// repoDigestFor returns the digest recorded for repo among an image's repo digests. +func repoDigestFor(repoDigests []string, repo string) string { + want := familiarRepo(repo) + for _, rd := range repoDigests { + name, digest, ok := strings.Cut(rd, "@") + if ok && familiarRepo(name) == want { + return digest + } + } + return "" +} + +// familiarRepo reduces a Docker Hub repository to the short form Docker prints ("library/nginx" -> "nginx"). +func familiarRepo(repo string) string { + repo = strings.TrimPrefix(repo, "docker.io/") + repo = strings.TrimPrefix(repo, "index.docker.io/") + if rest, ok := strings.CutPrefix(repo, "library/"); ok && !strings.Contains(rest, "/") { + return rest + } + return repo +} + func (sc *Scanner) isExcluded(image, tag string, exclusions []*UpdateExclusion) bool { + repo, _, _ := ParseImageRef(image) for _, e := range exclusions { switch e.PatternType { case ExclusionTypeImage: - if matched, _ := filepath.Match(e.Pattern, image); matched { - return true + // The full reference is still tried so that patterns written with a tag keep matching. + for _, candidate := range []string{repo, familiarRepo(repo), image} { + if matched, _ := filepath.Match(e.Pattern, candidate); matched { + return true + } } case ExclusionTypeTag: if matched, _ := filepath.Match(e.Pattern, tag); matched { diff --git a/internal/update/scanner_rules_test.go b/internal/update/scanner_rules_test.go new file mode 100644 index 00000000..9533f108 --- /dev/null +++ b/internal/update/scanner_rules_test.go @@ -0,0 +1,196 @@ +// Copyright 2026 Benjamin Touchard (kOlapsis) +// SPDX-License-Identifier: Apache-2.0 + +package update + +import ( + "context" + "testing" + + "github.com/stretchr/testify/assert" + "github.com/stretchr/testify/require" +) + +func oldBaseline(containerID, tag string) *DigestBaseline { + return &DigestBaseline{ContainerID: containerID, Tag: tag, RemoteDigest: "sha256:old"} +} + +func TestScanner_OfficialImageOnLatest_IsChecked(t *testing.T) { + for _, image := range []string{"nginx", "nginx:latest", "redis:latest"} { + t.Run(image, func(t *testing.T) { + reg := &stubRegistry{ + tags: map[string][]string{ + "library/nginx": {"latest", "1.27.0"}, + "library/redis": {"latest", "7.4.0"}, + }, + digest: "sha256:new", + } + sc := newTestScanner(reg, &stubStore{baseline: oldBaseline("ctr1", "latest")}) + + results, errs := sc.Scan(context.Background(), []ContainerInfo{ + {ExternalID: "ctr1", Name: "web", Image: image}, + }) + require.Empty(t, errs) + require.Len(t, results, 1) + assert.Equal(t, UpdateTypeDigestOnly, results[0].UpdateType) + assert.Equal(t, "latest", results[0].LatestTag) + }) + } +} + +func TestScanner_TrackLabels_LimitTheVersionJump(t *testing.T) { + cases := []struct { + name string + labels map[string]string + wantTag string + wantType UpdateType + }{ + {"no label", nil, "2.0.0", UpdateTypeMajor}, + {"track major", map[string]string{"maintenant.update.track": "major"}, "2.0.0", UpdateTypeMajor}, + {"track minor", map[string]string{"maintenant.update.track": "minor"}, "1.3.1", UpdateTypeMinor}, + {"track patch", map[string]string{"maintenant.update.track": "patch"}, "1.2.5", UpdateTypePatch}, + {"ignore_major", map[string]string{"maintenant.update.ignore_major": "true"}, "1.3.1", UpdateTypeMinor}, + {"ignore_major with track patch", map[string]string{ + "maintenant.update.ignore_major": "true", + "maintenant.update.track": "patch", + }, "1.2.5", UpdateTypePatch}, + } + for _, tc := range cases { + t.Run(tc.name, func(t *testing.T) { + reg := &stubRegistry{tags: map[string][]string{ + "library/postgres": {"1.2.3", "1.2.4", "1.2.5", "1.3.0", "1.3.1", "2.0.0"}, + }} + sc := newTestScanner(reg, &stubStore{}) + + results, errs := sc.Scan(context.Background(), []ContainerInfo{ + {ExternalID: "ctr1", Name: "db", Image: "postgres:1.2.3", Labels: tc.labels}, + }) + require.Empty(t, errs) + require.Len(t, results, 1) + assert.Equal(t, tc.wantTag, results[0].LatestTag) + assert.Equal(t, tc.wantType, results[0].UpdateType) + }) + } +} + +func TestScanner_TrackLabels_NothingInsideTheTrackedLevel(t *testing.T) { + reg := &stubRegistry{tags: map[string][]string{"library/postgres": {"1.2.3", "2.0.0"}}} + sc := newTestScanner(reg, &stubStore{}) + + results, errs := sc.Scan(context.Background(), []ContainerInfo{{ + ExternalID: "ctr1", Name: "db", Image: "postgres:1.2.3", + Labels: map[string]string{"maintenant.update.ignore_major": "true"}, + }}) + require.Empty(t, errs) + assert.Empty(t, results) +} + +// A floating tag whose next line is out of the tracked level still gets its digest compared. +func TestScanner_TrackMinor_FloatingTagFallsBackToDigest(t *testing.T) { + reg := &stubRegistry{ + tags: map[string][]string{"library/traefik": {"v3", "v3.7.10", "v4"}}, + digest: "sha256:new", + } + sc := newTestScanner(reg, &stubStore{baseline: oldBaseline("ctr1", "v3")}) + + results, errs := sc.Scan(context.Background(), []ContainerInfo{{ + ExternalID: "ctr1", Name: "proxy", Image: "traefik:v3", + Labels: map[string]string{"maintenant.update.track": "minor"}, + }}) + require.Empty(t, errs) + require.Len(t, results, 1) + assert.Equal(t, "v3", results[0].LatestTag) + assert.Equal(t, UpdateTypeDigestOnly, results[0].UpdateType) +} + +func TestScanner_DigestOnlyLabels_NeverListTags(t *testing.T) { + for _, labels := range []map[string]string{ + {"maintenant.update.digest_only": "true"}, + {"maintenant.update.track": "digest"}, + } { + reg := &stubRegistry{ + tags: map[string][]string{"library/postgres": {"16.4.0", "17.0.0"}}, + digest: "sha256:new", + } + sc := newTestScanner(reg, &stubStore{baseline: oldBaseline("ctr1", "16.4.0")}) + + results, errs := sc.Scan(context.Background(), []ContainerInfo{ + {ExternalID: "ctr1", Name: "db", Image: "postgres:16.4.0", Labels: labels}, + }) + require.Empty(t, errs) + require.Len(t, results, 1, "labels %v", labels) + assert.Equal(t, "16.4.0", results[0].LatestTag) + assert.Equal(t, UpdateTypeDigestOnly, results[0].UpdateType) + assert.Zero(t, reg.listCalls, "digest-only mode must not read the tag list") + } +} + +func TestScanner_DigestOnlyLabel_UnchangedDigestIsNoUpdate(t *testing.T) { + reg := &stubRegistry{digest: "sha256:old"} + sc := newTestScanner(reg, &stubStore{baseline: oldBaseline("ctr1", "16.4.0")}) + + results, errs := sc.Scan(context.Background(), []ContainerInfo{{ + ExternalID: "ctr1", Name: "db", Image: "postgres:16.4.0", + Labels: map[string]string{"maintenant.update.digest_only": "true"}, + }}) + require.Empty(t, errs) + assert.Empty(t, results) +} + +func TestScanner_ImageExclusion_ComparesTheNameWithoutTag(t *testing.T) { + cases := []struct { + pattern string + image string + excluded bool + }{ + {"myregistry.example.com/internal-app", "myregistry.example.com/internal-app:1.2.0", true}, + {"myregistry.example.com/internal-app", "myregistry.example.com/internal-app:1.2.0@sha256:abc", true}, + {"myregistry.example.com/*", "myregistry.example.com/internal-app:1.2.0", true}, + {"nginx", "docker.io/library/nginx:1.2.0", true}, + {"nginx:1.*", "nginx:1.2.0", true}, + {"myregistry.example.com/internal", "myregistry.example.com/internal-app:1.2.0", false}, + } + for _, tc := range cases { + t.Run(tc.pattern+" vs "+tc.image, func(t *testing.T) { + reg := &stubRegistry{tags: map[string][]string{ + "myregistry.example.com/internal-app": {"1.2.0", "1.3.0"}, + "library/nginx": {"1.2.0", "1.3.0"}, + }} + sc := newTestScanner(reg, &stubStore{exclusions: []*UpdateExclusion{ + {Pattern: tc.pattern, PatternType: ExclusionTypeImage}, + }}) + + results, errs := sc.Scan(context.Background(), []ContainerInfo{ + {ExternalID: "ctr1", Name: "app", Image: tc.image}, + }) + require.Empty(t, errs) + if tc.excluded { + assert.Empty(t, results) + } else { + assert.Len(t, results, 1) + } + }) + } +} + +// A tag filter that leaves out the current partial tag must not stop its digest comparison. +func TestScanner_PartialTag_FilterKeepsTheCurrentTagAsReference(t *testing.T) { + for _, labels := range []map[string]string{ + {"maintenant.update.tag-include": `^v3\.\d+\.\d+$`}, + {"maintenant.update.tag-exclude": `^v3$`}, + } { + reg := &stubRegistry{ + tags: map[string][]string{"library/traefik": {"v3", "v3.7.9", "v3.7.10"}}, + digest: "sha256:new", + } + sc := newTestScanner(reg, &stubStore{baseline: oldBaseline("ctr1", "v3")}) + + results, errs := sc.Scan(context.Background(), []ContainerInfo{ + {ExternalID: "ctr1", Name: "proxy", Image: "traefik:v3", Labels: labels}, + }) + require.Empty(t, errs) + require.Len(t, results, 1, "labels %v", labels) + assert.Equal(t, "v3", results[0].LatestTag) + assert.Equal(t, UpdateTypeDigestOnly, results[0].UpdateType) + } +} diff --git a/internal/update/scanner_test.go b/internal/update/scanner_test.go index 4ac01606..ead4577b 100644 --- a/internal/update/scanner_test.go +++ b/internal/update/scanner_test.go @@ -17,11 +17,13 @@ import ( // stubRegistry is a fake registryQuerier that returns predefined tag lists and digests. type stubRegistry struct { - tags map[string][]string // imageRef -> tags - digest string + tags map[string][]string // imageRef -> tags + digest string + listCalls int } func (r *stubRegistry) ListTags(_ context.Context, imageRef string) ([]string, error) { + r.listCalls++ if tags, ok := r.tags[imageRef]; ok { return tags, nil } @@ -34,7 +36,8 @@ func (r *stubRegistry) GetDigest(_ context.Context, _ string) (string, error) { // stubStore is a minimal UpdateStore that returns no pins, no exclusions, and no baseline. type stubStore struct { - baseline *DigestBaseline + baseline *DigestBaseline + exclusions []*UpdateExclusion } func (s *stubStore) InsertScanRecord(_ context.Context, _ *ScanRecord) (string, error) { @@ -78,8 +81,10 @@ func (s *stubStore) DeleteVersionPin(_ context.Context, _ string) error func (s *stubStore) InsertExclusion(_ context.Context, _ *UpdateExclusion) (string, error) { return "", nil } -func (s *stubStore) ListExclusions(_ context.Context) ([]*UpdateExclusion, error) { return nil, nil } -func (s *stubStore) DeleteExclusion(_ context.Context, _ string) error { return nil } +func (s *stubStore) ListExclusions(_ context.Context) ([]*UpdateExclusion, error) { + return s.exclusions, nil +} +func (s *stubStore) DeleteExclusion(_ context.Context, _ string) error { return nil } func (s *stubStore) InsertCVECacheEntry(_ context.Context, _ *CVECacheEntry) (string, error) { return "", nil } @@ -348,8 +353,6 @@ func TestScanner_InvalidTagExclude_FallsBackToDefault(t *testing.T) { // TestScanner_DigestOnlyMode_TagFilterBypassed verifies that tag filter labels are // ignored when the container uses a non-semver channel tag like "latest". // The digest comparison should proceed normally even if tag-include is set. -// Uses docker.io/library/nginx:latest to avoid the local-image heuristic check -// (which skips single-component images tagged "latest"). func TestScanner_DigestOnlyMode_TagFilterBypassed(t *testing.T) { reg := &stubRegistry{ tags: map[string][]string{ @@ -369,8 +372,6 @@ func TestScanner_DigestOnlyMode_TagFilterBypassed(t *testing.T) { sc := newTestScanner(reg, &stubStore{baseline: oldBaseline}) // tag-include set to semver pattern — should NOT affect "latest" digest comparison. - // Use docker.io/library/nginx:latest so imageRef is "library/nginx" (contains "/"), - // bypassing the local-image heuristic that skips bare "nginx:latest". containers := []ContainerInfo{ { ExternalID: "ctr1", diff --git a/internal/update/semver.go b/internal/update/semver.go index 77ab9d1c..0563c013 100644 --- a/internal/update/semver.go +++ b/internal/update/semver.go @@ -155,6 +155,15 @@ func SortTags(tags []string) []*semver.Version { return versions } +// isFixedVersionTag reports whether a tag names a single release (major.minor.patch) rather than a line the registry moves. +func isFixedVersionTag(tag string) bool { + versionPart, _ := splitVariant(tag) + if _, err := semver.NewVersion(versionPart); err != nil { + return false + } + return semverPrecision(versionPart) >= 3 +} + // digestOnly returns the current tag for digest comparison when the registry still // publishes it, so a republished or moved tag is detected. It never switches channel. func digestOnly(currentTag string, allTags []string) (string, UpdateType) { @@ -168,7 +177,7 @@ func digestOnly(currentTag string, allTags []string) (string, UpdateType) { // written the same way: same variant suffix, same "v" prefix, same number of numeric // components. Candidates whose major has far more digits than the current one are // build IDs (e.g. "608111629"), not releases. -func bestFloatingUpdate(currentVer *semver.Version, versionPart, variant string, allTags []string) *tagVersion { +func bestFloatingUpdate(currentVer *semver.Version, versionPart, variant string, allTags []string, level string) *tagVersion { candidates := sortTagVersions(allTags, variant, false, semverPrecision(versionPart)) var best *tagVersion @@ -181,13 +190,25 @@ func bestFloatingUpdate(currentVer *semver.Version, versionPart, variant string, if digitCount(c.version.Major()) > digitCount(currentVer.Major())+1 { continue } - if c.version.GreaterThan(currentVer) { + if c.version.GreaterThan(currentVer) && withinTrack(currentVer, c.version, level) { best = c } } return best } +// withinTrack reports whether moving from current to candidate stays inside the tracked level. +func withinTrack(current, candidate *semver.Version, level string) bool { + switch level { + case TrackMinor: + return candidate.Major() == current.Major() + case TrackPatch: + return candidate.Major() == current.Major() && candidate.Minor() == current.Minor() + default: + return true + } +} + // FindBestUpdate finds the best available update for the given current tag among all tags. // For pinned semver tags (major.minor.patch): finds the highest version with the same // variant suffix (e.g. -alpine). @@ -197,6 +218,12 @@ func bestFloatingUpdate(currentVer *semver.Version, versionPart, variant string, // Only a higher tag of the same shape counts; otherwise the same tag is returned so the // scanner compares digests. func FindBestUpdate(currentTag string, allTags []string) (bestTag string, updateType UpdateType) { + return findBestUpdate(currentTag, allTags, TrackMajor) +} + +// findBestUpdate is FindBestUpdate restricted to candidates within the tracked level +// (TrackMajor, TrackMinor or TrackPatch). +func findBestUpdate(currentTag string, allTags []string, level string) (bestTag string, updateType UpdateType) { versionPart, variant := splitVariant(currentTag) currentVer, err := semver.NewVersion(versionPart) @@ -207,7 +234,7 @@ func FindBestUpdate(currentTag string, allTags []string) (bestTag string, update // Partial version tag (e.g. "v3", "1.2", "16-bookworm"): floating channel. if semverPrecision(versionPart) < 3 { - best := bestFloatingUpdate(currentVer, versionPart, variant, allTags) + best := bestFloatingUpdate(currentVer, versionPart, variant, allTags, level) if best == nil { return digestOnly(currentTag, allTags) } @@ -224,7 +251,7 @@ func FindBestUpdate(currentTag string, allTags []string) (bestTag string, update // Find the highest version greater than current var best *tagVersion for i := range candidates { - if candidates[i].version.GreaterThan(currentVer) { + if candidates[i].version.GreaterThan(currentVer) && withinTrack(currentVer, candidates[i].version, level) { best = &candidates[i] } } diff --git a/internal/update/service.go b/internal/update/service.go index b41065ed..66019443 100644 --- a/internal/update/service.go +++ b/internal/update/service.go @@ -200,7 +200,7 @@ func (s *Service) ListImageUpdates(ctx context.Context, opts ListImageUpdatesOpt // GenerateUpdateCommand produces a shell command to update a container. func (s *Service) GenerateUpdateCommand(c ContainerInfo, latestTag string) string { - repo, _, _ := ParseImageRef(c.Image) + repo, currentTag, _ := ParseImageRef(c.Image) // Kubernetes workloads if c.RuntimeType == "kubernetes" && c.ControllerKind != "" { @@ -210,13 +210,14 @@ func (s *Service) GenerateUpdateCommand(c ContainerInfo, latestTag string) strin } // Docker Compose - if c.RuntimeType != "kubernetes" && c.OrchestrationGroup != "" && c.OrchestrationUnit != "" { - dir := c.ComposeWorkingDir - if dir == "" { - dir = "" + if isCompose(c) { + // Compose pulls the tag written in the compose file, so a new tag has to be written there first. + if latestTag != currentTag { + return fmt.Sprintf("cd %s\n# Set the image of service %s to %s:%s in the compose file, then:\ndocker compose pull %s\ndocker compose up -d %s", + composeDir(c), c.OrchestrationUnit, repo, latestTag, c.OrchestrationUnit, c.OrchestrationUnit) } return fmt.Sprintf("cd %s\ndocker compose pull %s\ndocker compose up -d --force-recreate %s", - dir, c.OrchestrationUnit, c.OrchestrationUnit) + composeDir(c), c.OrchestrationUnit, c.OrchestrationUnit) } // Standalone Docker container @@ -224,34 +225,66 @@ func (s *Service) GenerateUpdateCommand(c ContainerInfo, latestTag string) strin repo, latestTag, c.Name, c.Name, c.Name, repo, latestTag) } -// GenerateRollbackCommand produces a shell command to revert a container to its previous image digest. -func (s *Service) GenerateRollbackCommand(c ContainerInfo, previousDigest string) string { - if previousDigest == "" { +// GenerateRollbackCommand produces a shell command that puts a container back on the image it ran before the update, or "" when that image cannot be named. +func (s *Service) GenerateRollbackCommand(c ContainerInfo, u *ImageUpdate) string { + ref := previousImageRef(u.Image, u.CurrentTag, u.PreviousDigest) + if ref == "" { return "" } - repo, _, _ := ParseImageRef(c.Image) - - // Kubernetes workloads — use rollout undo + // Kubernetes workloads if c.RuntimeType == "kubernetes" && c.ControllerKind != "" { kind := strings.ToLower(c.ControllerKind) - return fmt.Sprintf("kubectl rollout undo %s/%s -n %s", - kind, c.OrchestrationUnit, c.OrchestrationGroup) + return fmt.Sprintf("kubectl set image %s/%s %s=%s -n %s", + kind, c.OrchestrationUnit, c.Name, ref, c.OrchestrationGroup) } - // Docker Compose — recreate with previous digest - if c.RuntimeType != "kubernetes" && c.OrchestrationGroup != "" && c.OrchestrationUnit != "" { - dir := c.ComposeWorkingDir - if dir == "" { - dir = "" + // Docker Compose + if isCompose(c) { + svc := c.OrchestrationUnit + if u.LatestTag == u.CurrentTag { + // The compose file keeps the same tag: point that tag back at the previous image locally. + return fmt.Sprintf("cd %s\ndocker pull %s\ndocker tag %s %s\ndocker compose up -d --pull never --force-recreate %s", + composeDir(c), ref, ref, imageWithoutDigest(u.Image), svc) } - return fmt.Sprintf("cd %s\ndocker compose pull %s\ndocker compose up -d --force-recreate %s", - dir, c.OrchestrationUnit, c.OrchestrationUnit) + return fmt.Sprintf("cd %s\n# Set the image of service %s back to %s in the compose file, then:\ndocker compose up -d %s", + composeDir(c), svc, ref, svc) } - // Standalone Docker container — stop/rm/run with digest reference - return fmt.Sprintf("docker stop %s && docker rm %s\ndocker run -d --name %s %s@%s", - c.Name, c.Name, c.Name, repo, previousDigest) + // Standalone Docker container + return fmt.Sprintf("docker pull %s\ndocker stop %s && docker rm %s\ndocker run -d --name %s %s", + ref, c.Name, c.Name, c.Name, ref) +} + +// previousImageRef names the image a container ran before an update: by digest when known, +// else by its tag when that tag is a fixed release; a moving tag without digest cannot name it. +func previousImageRef(image, currentTag, previousDigest string) string { + repo, _, _ := ParseImageRef(image) + if previousDigest != "" { + return repo + "@" + previousDigest + } + if isFixedVersionTag(currentTag) { + return repo + ":" + currentTag + } + return "" +} + +func imageWithoutDigest(image string) string { + if i := strings.Index(image, "@"); i > 0 { + return image[:i] + } + return image +} + +func isCompose(c ContainerInfo) bool { + return c.RuntimeType != "kubernetes" && c.OrchestrationGroup != "" && c.OrchestrationUnit != "" +} + +func composeDir(c ContainerInfo) string { + if c.ComposeWorkingDir == "" { + return "" + } + return c.ComposeWorkingDir } // GenerateFixCommand produces a shell command to update a container to a specific CVE fix version. @@ -371,19 +404,20 @@ func (s *Service) runScan(ctx context.Context) { riskScore := BaseRiskScore(r.UpdateType) u := &ImageUpdate{ - ScanID: scanID, - ContainerID: r.ContainerID, - ContainerName: r.ContainerName, - Image: r.Image, - CurrentTag: r.CurrentTag, - CurrentDigest: r.CurrentDigest, - Registry: r.Registry, - LatestTag: r.LatestTag, - LatestDigest: r.LatestDigest, - UpdateType: r.UpdateType, - RiskScore: riskScore, - Status: StatusAvailable, - DetectedAt: time.Now(), + ScanID: scanID, + ContainerID: r.ContainerID, + ContainerName: r.ContainerName, + Image: r.Image, + CurrentTag: r.CurrentTag, + CurrentDigest: r.CurrentDigest, + Registry: r.Registry, + LatestTag: r.LatestTag, + LatestDigest: r.LatestDigest, + UpdateType: r.UpdateType, + RiskScore: riskScore, + PreviousDigest: r.PreviousDigest, + Status: StatusAvailable, + DetectedAt: time.Now(), } if _, err := s.store.InsertImageUpdate(ctx, u); err != nil { @@ -402,12 +436,13 @@ func (s *Service) runScan(ctx context.Context) { "latest_tag": r.LatestTag, "update_type": string(r.UpdateType), "risk_score": riskScore, + "alert_on": r.AlertOn, } if ci, ok := containerByID[r.ContainerID]; ok { eventData["update_command"] = s.GenerateUpdateCommand(ci, r.LatestTag) - if r.CurrentDigest != "" { - eventData["rollback_command"] = s.GenerateRollbackCommand(ci, r.CurrentDigest) + if cmd := s.GenerateRollbackCommand(ci, u); cmd != "" { + eventData["rollback_command"] = cmd } } diff --git a/internal/update/service_commands_test.go b/internal/update/service_commands_test.go new file mode 100644 index 00000000..16729502 --- /dev/null +++ b/internal/update/service_commands_test.go @@ -0,0 +1,205 @@ +// Copyright 2026 Benjamin Touchard (kOlapsis) +// SPDX-License-Identifier: Apache-2.0 + +package update + +import ( + "context" + "testing" + + "github.com/stretchr/testify/assert" + "github.com/stretchr/testify/require" + + "github.com/kolapsis/maintenant/internal/event" +) + +func composeWeb(image string) ContainerInfo { + return ContainerInfo{ + Name: "app-web-1", + Image: image, + OrchestrationGroup: "app", + OrchestrationUnit: "web", + ComposeWorkingDir: "/srv/app", + RuntimeType: "docker", + } +} + +func TestGenerateRollbackCommand_ComposeSameTag_RetagsThePreviousImage(t *testing.T) { + svc := &Service{} + c := composeWeb("nginx:latest") + u := &ImageUpdate{Image: "nginx:latest", CurrentTag: "latest", LatestTag: "latest", PreviousDigest: "sha256:old"} + + got := svc.GenerateRollbackCommand(c, u) + + assert.Equal(t, "cd /srv/app\n"+ + "docker pull nginx@sha256:old\n"+ + "docker tag nginx@sha256:old nginx:latest\n"+ + "docker compose up -d --pull never --force-recreate web", got) + assert.NotEqual(t, svc.GenerateUpdateCommand(c, "latest"), got) +} + +func TestGenerateRollbackCommand_ComposeNewTag_PinsThePreviousImage(t *testing.T) { + svc := &Service{} + c := composeWeb("nginx:1.24.0") + + byDigest := svc.GenerateRollbackCommand(c, &ImageUpdate{ + Image: "nginx:1.24.0", CurrentTag: "1.24.0", LatestTag: "1.26.0", PreviousDigest: "sha256:old", + }) + assert.Equal(t, "cd /srv/app\n"+ + "# Set the image of service web back to nginx@sha256:old in the compose file, then:\n"+ + "docker compose up -d web", byDigest) + + byTag := svc.GenerateRollbackCommand(c, &ImageUpdate{ + Image: "nginx:1.24.0", CurrentTag: "1.24.0", LatestTag: "1.26.0", + }) + assert.Contains(t, byTag, "back to nginx:1.24.0 in the compose file") +} + +func TestGenerateRollbackCommand_Standalone(t *testing.T) { + svc := &Service{} + c := ContainerInfo{Name: "web", Image: "ghcr.io/acme/web:2.1.0", RuntimeType: "docker"} + + got := svc.GenerateRollbackCommand(c, &ImageUpdate{ + Image: "ghcr.io/acme/web:2.1.0", CurrentTag: "2.1.0", LatestTag: "2.2.0", PreviousDigest: "sha256:old", + }) + + assert.Equal(t, "docker pull ghcr.io/acme/web@sha256:old\n"+ + "docker stop web && docker rm web\n"+ + "docker run -d --name web ghcr.io/acme/web@sha256:old", got) +} + +func TestGenerateRollbackCommand_Kubernetes(t *testing.T) { + svc := &Service{} + c := ContainerInfo{ + Name: "api", Image: "ghcr.io/acme/api:1.0.0", RuntimeType: "kubernetes", + ControllerKind: "Deployment", OrchestrationUnit: "api", OrchestrationGroup: "prod", + } + + got := svc.GenerateRollbackCommand(c, &ImageUpdate{ + Image: "ghcr.io/acme/api:1.0.0", CurrentTag: "1.0.0", LatestTag: "1.1.0", + }) + + assert.Equal(t, "kubectl set image deployment/api api=ghcr.io/acme/api:1.0.0 -n prod", got) +} + +// A moving tag without a known digest cannot name the image it pointed to before. +func TestGenerateRollbackCommand_MovingTagWithoutDigest(t *testing.T) { + svc := &Service{} + for _, tag := range []string{"latest", "v3", "1.2"} { + got := svc.GenerateRollbackCommand(composeWeb("traefik:"+tag), &ImageUpdate{ + Image: "traefik:" + tag, CurrentTag: tag, LatestTag: tag, + }) + assert.Empty(t, got, tag) + } +} + +func TestGenerateUpdateCommand_ComposeNewTagIsWrittenInTheComposeFile(t *testing.T) { + svc := &Service{} + + assert.Equal(t, "cd /srv/app\n"+ + "# Set the image of service web to nginx:1.26.0 in the compose file, then:\n"+ + "docker compose pull web\n"+ + "docker compose up -d web", svc.GenerateUpdateCommand(composeWeb("nginx:1.24.0"), "1.26.0")) + + assert.Equal(t, "cd /srv/app\n"+ + "docker compose pull web\n"+ + "docker compose up -d --force-recreate web", svc.GenerateUpdateCommand(composeWeb("nginx:latest"), "latest")) +} + +func TestScanner_LocallyBuiltImage_IsSkipped(t *testing.T) { + for _, image := range []string{"myapp:latest", "ghcr.io/acme/app:1.0.0"} { + reg := &stubRegistry{ + tags: map[string][]string{ + "library/myapp": {"latest"}, + "ghcr.io/acme/app": {"1.0.0", "2.0.0"}, + }, + digest: "sha256:new", + } + sc := newTestScanner(reg, &stubStore{baseline: oldBaseline("ctr1", "latest")}) + + results, errs := sc.Scan(context.Background(), []ContainerInfo{ + {ExternalID: "ctr1", Name: "app", Image: image, LocallyBuilt: true}, + }) + require.Empty(t, errs) + assert.Empty(t, results, image) + assert.Zero(t, reg.listCalls, "a locally built image has no registry to ask") + } +} + +func TestScanner_PreviousDigestIsTheRunningImage(t *testing.T) { + reg := &stubRegistry{ + tags: map[string][]string{"library/nginx": {"latest"}}, + digest: "sha256:new", + } + sc := newTestScanner(reg, &stubStore{baseline: oldBaseline("ctr1", "latest")}) + + results, errs := sc.Scan(context.Background(), []ContainerInfo{{ + ExternalID: "ctr1", Name: "web", Image: "nginx:latest", + RepoDigests: []string{"mirror.example.com/nginx@sha256:mirror", "docker.io/library/nginx@sha256:running"}, + }}) + require.Empty(t, errs) + require.Len(t, results, 1) + assert.Equal(t, "sha256:running", results[0].PreviousDigest) +} + +// captureStore records the updates a scan persists. +type captureStore struct { + *stubStore + inserted []*ImageUpdate +} + +func (s *captureStore) InsertImageUpdate(_ context.Context, u *ImageUpdate) (string, error) { + s.inserted = append(s.inserted, u) + return "", nil +} + +func TestRunScan_PersistsThePreviousImageAndCarriesAlertOn(t *testing.T) { + store := &captureStore{stubStore: &stubStore{}} + reg := &stubRegistry{ + tags: map[string][]string{"library/nginx": {"1.24.0", "1.26.0"}}, + digest: "sha256:latest", + } + + var detected map[string]interface{} + svc := NewService(Deps{ + Store: store, + Scanner: newTestScanner(reg, store), + Containers: stubLister{containers: []ContainerInfo{{ + UID: "uid1", ExternalID: "ctr1", Name: "web", Image: "nginx:1.24.0", + RepoDigests: []string{"nginx@sha256:running"}, + Labels: map[string]string{"maintenant.update.alert_on": "critical"}, + }}}, + Logger: testLogger(), + EventCallback: func(eventType string, data interface{}) { + if eventType == event.UpdateDetected { + detected, _ = data.(map[string]interface{}) + } + }, + }) + + svc.runScan(context.Background()) + + require.Len(t, store.inserted, 1) + assert.Equal(t, "sha256:running", store.inserted[0].PreviousDigest) + require.NotNil(t, detected) + assert.Equal(t, AlertOnCritical, detected["alert_on"]) + assert.Equal(t, "docker pull nginx@sha256:running\n"+ + "docker stop web && docker rm web\n"+ + "docker run -d --name web nginx@sha256:running", detected["rollback_command"]) +} + +func TestUpdateConfig_TrackLevel(t *testing.T) { + cases := []struct { + cfg UpdateConfig + want string + }{ + {UpdateConfig{}, TrackMajor}, + {UpdateConfig{IgnoreMajor: true}, TrackMinor}, + {UpdateConfig{Track: TrackPatch, IgnoreMajor: true}, TrackPatch}, + {UpdateConfig{Track: TrackMinor, DigestOnly: true}, TrackDigest}, + {UpdateConfig{Track: TrackDigest}, TrackDigest}, + } + for _, tc := range cases { + assert.Equal(t, tc.want, tc.cfg.TrackLevel(), "%+v", tc.cfg) + } +} From 6b4492c462c72cd8e60d9594d3bfefb953341f37 Mon Sep 17 00:00:00 2001 From: Benjamin Date: Wed, 30 Sep 2026 16:23:29 +0200 Subject: [PATCH 03/54] fix(http): refuse cross-origin API writes, unbuffer SSE, explain TLS refusals Unsafe cross-origin browser requests on /api/ are refused with a 403 CROSS_ORIGIN_REFUSED through http.CrossOriginProtection, trusting the origins listed in MAINTENANT_CORS_ORIGINS. /ping, /status, /mcp, /oauth and /.well-known keep accepting cross-origin calls. CORS preflights now allow PATCH, every SSE response (MCP transport included) carries X-Accel-Buffering: no, a PostgreSQL server that declines TLS is reported apart from an unreachable one, and a MCP client secret shorter than 32 characters logs a warning at startup. --- cmd/maintenant/storage_errors.go | 6 + cmd/maintenant/storage_errors_test.go | 20 +++ internal/api/v1/cross_origin_test.go | 183 ++++++++++++++++++++++++++ internal/api/v1/logs_stream.go | 8 +- internal/api/v1/logs_stream_test.go | 2 +- internal/api/v1/middleware.go | 30 ++++- internal/api/v1/router.go | 8 +- internal/api/v1/sse.go | 12 +- internal/api/v1/sse_headers_test.go | 72 ++++++++++ internal/app/http.go | 29 +++- internal/app/http_wiring_test.go | 171 ++++++++++++++++++++++++ internal/store/db_postgres.go | 65 ++++++++- internal/store/db_postgres_test.go | 61 +++++++++ internal/store/dsn.go | 16 ++- internal/store/dsn_test.go | 16 +++ internal/store/errors.go | 2 + internal/store/errors_test.go | 2 +- 17 files changed, 677 insertions(+), 26 deletions(-) create mode 100644 internal/api/v1/cross_origin_test.go create mode 100644 internal/api/v1/sse_headers_test.go create mode 100644 internal/app/http_wiring_test.go diff --git a/cmd/maintenant/storage_errors.go b/cmd/maintenant/storage_errors.go index a08c843e..aa59573d 100644 --- a/cmd/maintenant/storage_errors.go +++ b/cmd/maintenant/storage_errors.go @@ -24,6 +24,12 @@ func logStorageStartupError(logger *slog.Logger, err error, dsn string) bool { case errors.Is(err, store.ErrInvalidDSN): logger.Error("the database connection string cannot be read", "fix", "expected postgres://user:password@host:5432/database[?sslmode=require]") + case errors.Is(err, store.ErrTLSRefused) && store.DefaultsSSLMode(dsn): + logger.Error("the database server does not accept TLS", "target", target, + "fix", "sslmode=require is added by default for a non-local host: enable TLS on the PostgreSQL server, or set sslmode=disable explicitly in MAINTENANT_DATABASE_URL if the network to the database is trusted") + case errors.Is(err, store.ErrTLSRefused): + logger.Error("the database server does not accept TLS", "target", target, + "fix", "the connection string requires TLS through its sslmode: enable TLS on the PostgreSQL server, or lower sslmode if the network to the database is trusted") case errors.Is(err, store.ErrUnreachable): logger.Error("the database does not answer", "target", target, "fix", "check the host and port, the network route and the firewall; the instance does not fall back to the local file") diff --git a/cmd/maintenant/storage_errors_test.go b/cmd/maintenant/storage_errors_test.go index b714bfcb..7e45d363 100644 --- a/cmd/maintenant/storage_errors_test.go +++ b/cmd/maintenant/storage_errors_test.go @@ -31,6 +31,7 @@ func TestLogStorageStartupError_DistinctFamilies(t *testing.T) { "agent mode": fmt.Errorf("wrapped: %w", app.ErrDatabaseURLInAgentMode), "invalid dsn": fmt.Errorf("open database: %w", store.ErrInvalidDSN), "unreachable": fmt.Errorf("open database: %w", store.ErrUnreachable), + "tls refused": fmt.Errorf("open database: %w", store.ErrTLSRefused), "credentials": fmt.Errorf("open database: %w", store.ErrAuthRefused), "version": fmt.Errorf("open database: %w", store.ErrUnsupportedVersion), "schema from future": fmt.Errorf("run migrations: %w", store.ErrSchemaNewer), @@ -64,6 +65,25 @@ func TestLogStorageStartupError_DistinctFamilies(t *testing.T) { assert.Contains(t, messages["version"], "14") } +func TestLogStorageStartupError_TLSRefusedNamesTheDefault(t *testing.T) { + err := fmt.Errorf("open database: %w", store.ErrTLSRefused) + + var defaulted bytes.Buffer + require.True(t, logStorageStartupError(slog.New(slog.NewTextHandler(&defaulted, nil)), err, + "postgres://app:"+cliSentinelPassword+"@db:5432/maintenant")) + assert.Contains(t, defaulted.String(), "does not accept TLS") + assert.Contains(t, defaulted.String(), "sslmode=require is added by default") + assert.Contains(t, defaulted.String(), "sslmode=disable") + assert.NotContains(t, defaulted.String(), "firewall") + assert.NotContains(t, defaulted.String(), cliSentinelPassword) + + var explicit bytes.Buffer + require.True(t, logStorageStartupError(slog.New(slog.NewTextHandler(&explicit, nil)), err, + "postgres://app@db:5432/maintenant?sslmode=require")) + assert.Contains(t, explicit.String(), "does not accept TLS") + assert.NotContains(t, explicit.String(), "added by default", "the operator asked for TLS") +} + // TestLogStorageStartupError_PassesThroughOtherErrors keeps the classifier // honest: an error that is not a storage startup failure is left to the // caller's generic handler rather than mislabelled. diff --git a/internal/api/v1/cross_origin_test.go b/internal/api/v1/cross_origin_test.go new file mode 100644 index 00000000..ac59027e --- /dev/null +++ b/internal/api/v1/cross_origin_test.go @@ -0,0 +1,183 @@ +// Copyright 2026 Benjamin Touchard (kOlapsis) +// SPDX-License-Identifier: Apache-2.0 + +package v1 + +import ( + "bytes" + "encoding/json" + "io" + "log/slog" + "net/http" + "net/http/httptest" + "strings" + "testing" + + "github.com/stretchr/testify/assert" + "github.com/stretchr/testify/require" +) + +const webhookBody = `{"name":"hook","url":"https://8.8.8.8/hook","event_types":["*"]}` + +func webhookRouter(t *testing.T, corsOrigins string, logs io.Writer) (*Router, *stubWebhookStore) { + t.Helper() + if logs == nil { + logs = io.Discard + } + store := &stubWebhookStore{} + r := NewRouter(HandlerDeps{ + Logger: slog.New(slog.NewTextHandler(logs, nil)), + WebhookStore: store, + CORSOrigins: corsOrigins, + }) + return r, store +} + +func sendWebhook(r *Router, method string, headers map[string]string) *httptest.ResponseRecorder { + var body io.Reader + if method != http.MethodGet { + body = strings.NewReader(webhookBody) + } + req := httptest.NewRequest(method, "http://maintenant.example.com/api/v1/webhooks", body) + for k, v := range headers { + req.Header.Set(k, v) + } + rec := httptest.NewRecorder() + r.Handler().ServeHTTP(rec, req) + return rec +} + +func TestCrossOrigin_RefusesCrossSiteWrite(t *testing.T) { + for name, headers := range map[string]map[string]string{ + "modern browser": { + "Sec-Fetch-Site": "cross-site", + "Origin": "https://evil.example", + "Content-Type": "text/plain", + }, + "same site, other origin": { + "Sec-Fetch-Site": "same-site", + "Origin": "https://blog.example.com", + "Content-Type": "text/plain", + }, + "browser without Sec-Fetch-Site": { + "Origin": "https://evil.example", + "Content-Type": "text/plain", + }, + } { + t.Run(name, func(t *testing.T) { + r, store := webhookRouter(t, "", nil) + + rec := sendWebhook(r, http.MethodPost, headers) + + require.Equal(t, http.StatusForbidden, rec.Code) + assert.Nil(t, store.created, "a refused request must never reach the handler") + assert.Equal(t, "application/json", rec.Header().Get("Content-Type")) + var body ErrorResponse + require.NoError(t, json.Unmarshal(rec.Body.Bytes(), &body)) + assert.Equal(t, "CROSS_ORIGIN_REFUSED", body.Error.Code) + assert.Contains(t, body.Error.Message, "MAINTENANT_CORS_ORIGINS") + }) + } +} + +func TestCrossOrigin_AllowsSameOriginAndNonBrowserWrites(t *testing.T) { + for name, headers := range map[string]map[string]string{ + "same origin": {"Sec-Fetch-Site": "same-origin", "Origin": "http://maintenant.example.com"}, + "user initiated": {"Sec-Fetch-Site": "none"}, + "old browser, Origin is Host": {"Origin": "http://maintenant.example.com"}, + "script without Origin (curl)": {}, + } { + t.Run(name, func(t *testing.T) { + r, store := webhookRouter(t, "", nil) + + rec := sendWebhook(r, http.MethodPost, headers) + + assert.Equal(t, http.StatusCreated, rec.Code, rec.Body.String()) + assert.NotNil(t, store.created) + }) + } +} + +func TestCrossOrigin_TrustsListedCORSOrigins(t *testing.T) { + r, store := webhookRouter(t, "https://ops.example.org, https://grafana.example.org", nil) + + rec := sendWebhook(r, http.MethodPost, map[string]string{ + "Sec-Fetch-Site": "cross-site", + "Origin": "https://grafana.example.org", + "Content-Type": "application/json", + }) + assert.Equal(t, http.StatusCreated, rec.Code, rec.Body.String()) + assert.NotNil(t, store.created) + + rec = sendWebhook(r, http.MethodPost, map[string]string{ + "Sec-Fetch-Site": "cross-site", + "Origin": "https://evil.example", + }) + assert.Equal(t, http.StatusForbidden, rec.Code, "only the listed origins are trusted") +} + +func TestCrossOrigin_WildcardCORSTrustsNoOrigin(t *testing.T) { + r, store := webhookRouter(t, "*", nil) + + rec := sendWebhook(r, http.MethodPost, map[string]string{ + "Sec-Fetch-Site": "cross-site", + "Origin": "https://evil.example", + }) + + assert.Equal(t, http.StatusForbidden, rec.Code) + assert.Nil(t, store.created) +} + +func TestCrossOrigin_NeverBlocksSafeMethods(t *testing.T) { + r, _ := webhookRouter(t, "", nil) + + rec := sendWebhook(r, http.MethodGet, map[string]string{ + "Sec-Fetch-Site": "cross-site", + "Origin": "https://evil.example", + }) + + assert.Equal(t, http.StatusOK, rec.Code) +} + +func TestCrossOrigin_LeavesPingRoutesOpen(t *testing.T) { + var reached bool + h := crossOriginGuard(http.NewCrossOriginProtection(), http.HandlerFunc(func(w http.ResponseWriter, _ *http.Request) { + reached = true + w.WriteHeader(http.StatusOK) + })) + + req := httptest.NewRequest(http.MethodPost, "/ping/0b7c3a4e-7d1f-4c62-9a55-3f3b8d0c2e11", nil) + req.Header.Set("Sec-Fetch-Site", "cross-site") + req.Header.Set("Origin", "https://ci.example") + rec := httptest.NewRecorder() + h.ServeHTTP(rec, req) + + assert.True(t, reached, "a heartbeat ping may come from anywhere") + assert.Equal(t, http.StatusOK, rec.Code) +} + +func TestCrossOrigin_WarnsAboutAnEntryThatIsNotAnOrigin(t *testing.T) { + var logs bytes.Buffer + r, _ := webhookRouter(t, "https://ops.example.org/, https://grafana.example.org", &logs) + + assert.Contains(t, logs.String(), "level=WARN") + assert.Contains(t, logs.String(), "https://ops.example.org/") + + rec := sendWebhook(r, http.MethodPost, map[string]string{ + "Sec-Fetch-Site": "cross-site", + "Origin": "https://grafana.example.org", + }) + assert.Equal(t, http.StatusCreated, rec.Code, "the valid entries are still trusted") +} + +func TestCORS_AllowsPatch(t *testing.T) { + handler := cors([]string{"https://ops.example.org"}, http.HandlerFunc(func(w http.ResponseWriter, _ *http.Request) {})) + + req := httptest.NewRequest(http.MethodOptions, "/api/v1/agents/1", nil) + req.Header.Set("Origin", "https://ops.example.org") + req.Header.Set("Access-Control-Request-Method", http.MethodPatch) + rec := httptest.NewRecorder() + handler.ServeHTTP(rec, req) + + assert.Contains(t, rec.Header().Get("Access-Control-Allow-Methods"), "PATCH") +} diff --git a/internal/api/v1/logs_stream.go b/internal/api/v1/logs_stream.go index 37f84d27..4f3bf3eb 100644 --- a/internal/api/v1/logs_stream.go +++ b/internal/api/v1/logs_stream.go @@ -139,9 +139,7 @@ func (h *LogStreamHandler) HandleLogStream(w http.ResponseWriter, r *http.Reques // Set SSE headers. CORS is deliberately absent: the cors() middleware // applies the configured policy, and forcing a wildcard here let any site // an operator visited read this container's logs. - w.Header().Set("Content-Type", "text/event-stream") - w.Header().Set("Cache-Control", "no-cache") - w.Header().Set("Connection", "keep-alive") + setSSEHeaders(w.Header()) flusher.Flush() scanner := bufio.NewScanner(reader) @@ -232,9 +230,7 @@ func (h *LogStreamHandler) streamRemote( defer release() // Same as the local path: no wildcard CORS, cors() owns the policy. - w.Header().Set("Content-Type", "text/event-stream") - w.Header().Set("Cache-Control", "no-cache") - w.Header().Set("Connection", "keep-alive") + setSSEHeaders(w.Header()) flusher.Flush() ctx := r.Context() diff --git a/internal/api/v1/logs_stream_test.go b/internal/api/v1/logs_stream_test.go index cb380b0a..5c83ff0d 100644 --- a/internal/api/v1/logs_stream_test.go +++ b/internal/api/v1/logs_stream_test.go @@ -110,7 +110,7 @@ func TestHandleLogStream(t *testing.T) { assert.Equal(t, tt.wantStatus, w.Code) if tt.wantSSE { - assert.Equal(t, "text/event-stream", w.Header().Get("Content-Type")) + assertUnbufferedEventStream(t, w.Header()) assert.Contains(t, w.Body.String(), tt.wantContains) } }) diff --git a/internal/api/v1/middleware.go b/internal/api/v1/middleware.go index 8ebcb232..239c7754 100644 --- a/internal/api/v1/middleware.go +++ b/internal/api/v1/middleware.go @@ -133,7 +133,7 @@ func cors(origins []string, next http.Handler) http.Handler { } } } - w.Header().Set("Access-Control-Allow-Methods", "GET, POST, PUT, DELETE, OPTIONS") + w.Header().Set("Access-Control-Allow-Methods", "GET, POST, PUT, PATCH, DELETE, OPTIONS") w.Header().Set("Access-Control-Allow-Headers", "Content-Type, Authorization") } @@ -146,6 +146,34 @@ func cors(origins []string, next http.Handler) http.Handler { }) } +const crossOriginMessage = "Cross-origin request refused: add the calling origin to MAINTENANT_CORS_ORIGINS to allow it." + +// newCrossOriginProtection trusts the origins MAINTENANT_CORS_ORIGINS lists, and none for the wildcard. +func newCrossOriginProtection(origins []string, logger *slog.Logger) *http.CrossOriginProtection { + cop := http.NewCrossOriginProtection() + if len(origins) == 1 && origins[0] == "*" { + return cop + } + for _, origin := range origins { + if err := cop.AddTrustedOrigin(origin); err != nil { + logger.Warn("MAINTENANT_CORS_ORIGINS entry is not an origin and matches no request", + "entry", origin, "expected", "scheme://host[:port]", "error", err) + } + } + return cop +} + +// crossOriginGuard refuses a browser request with an unsafe method on the API when it comes from another, untrusted origin. +func crossOriginGuard(cop *http.CrossOriginProtection, next http.Handler) http.Handler { + return http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) { + if strings.HasPrefix(r.URL.Path, "/api/") && cop.Check(r) != nil { + WriteError(w, http.StatusForbidden, "CROSS_ORIGIN_REFUSED", crossOriginMessage) + return + } + next.ServeHTTP(w, r) + }) +} + // bodyLimit caps the request body size, whatever the method. func bodyLimit(maxBytes int64, next http.Handler) http.Handler { return http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) { diff --git a/internal/api/v1/router.go b/internal/api/v1/router.go index a81bf1bd..5dedb731 100644 --- a/internal/api/v1/router.go +++ b/internal/api/v1/router.go @@ -178,6 +178,7 @@ type Router struct { storage StorageStatus containerHandler *ContainerHandler corsOrigins []string + crossOrigin *http.CrossOriginProtection maxBodySize int64 buildVersion string organisationName string @@ -192,13 +193,15 @@ func NewRouter(d HandlerDeps) *Router { maxBody = 1048576 // 1 MB default } + corsOrigins := parseCORSOrigins(d.CORSOrigins) r := &Router{ mux: http.NewServeMux(), broker: d.Broker, logger: d.Logger, runtime: d.Runtime, storage: d.Storage, - corsOrigins: parseCORSOrigins(d.CORSOrigins), + corsOrigins: corsOrigins, + crossOrigin: newCrossOriginProtection(corsOrigins, d.Logger), maxBodySize: maxBody, buildVersion: d.BuildVersion, organisationName: d.OrganisationName, @@ -737,10 +740,11 @@ func (r *Router) registerSwarmRoutes(d HandlerDeps) { } // Handler returns the HTTP handler with the full middleware chain applied. -// Middleware order (outermost to innermost): panicRecovery → requestLogger → requestID → cors → bodyLimit → mux +// Middleware order (outermost to innermost): panicRecovery → requestLogger → requestID → cors → crossOriginGuard → bodyLimit → mux func (r *Router) Handler() http.Handler { var h http.Handler = r.mux h = bodyLimit(r.maxBodySize, h) + h = crossOriginGuard(r.crossOrigin, h) h = cors(r.corsOrigins, h) h = requestID(h) h = requestLogger(h, r.logger) diff --git a/internal/api/v1/sse.go b/internal/api/v1/sse.go index bc96ad28..49d2fe1b 100644 --- a/internal/api/v1/sse.go +++ b/internal/api/v1/sse.go @@ -84,9 +84,7 @@ func (b *SSEBroker) ServeHTTP(w http.ResponseWriter, r *http.Request) { return } - w.Header().Set("Content-Type", "text/event-stream") - w.Header().Set("Cache-Control", "no-cache") - w.Header().Set("Connection", "keep-alive") + setSSEHeaders(w.Header()) ch := make(chan SSEEvent, 64) @@ -119,6 +117,14 @@ func (b *SSEBroker) ServeHTTP(w http.ResponseWriter, r *http.Request) { } } +// setSSEHeaders marks a response as an event stream that no cache or buffering proxy may hold back. +func setSSEHeaders(h http.Header) { + h.Set("Content-Type", "text/event-stream") + h.Set("Cache-Control", "no-cache") + h.Set("Connection", "keep-alive") + h.Set("X-Accel-Buffering", "no") +} + // ClientCount returns the number of connected SSE clients. func (b *SSEBroker) ClientCount() int { b.mu.RLock() diff --git a/internal/api/v1/sse_headers_test.go b/internal/api/v1/sse_headers_test.go new file mode 100644 index 00000000..4e305b50 --- /dev/null +++ b/internal/api/v1/sse_headers_test.go @@ -0,0 +1,72 @@ +// Copyright 2026 Benjamin Touchard (kOlapsis) +// SPDX-License-Identifier: Apache-2.0 + +package v1 + +import ( + "context" + "io" + "log/slog" + "net/http" + "net/http/httptest" + "testing" + + "github.com/stretchr/testify/assert" + + "github.com/kolapsis/maintenant/internal/agentpb" + "github.com/kolapsis/maintenant/internal/container" +) + +func assertUnbufferedEventStream(t *testing.T, h http.Header) { + t.Helper() + assert.Equal(t, "text/event-stream", h.Get("Content-Type")) + assert.Equal(t, "no", h.Get("X-Accel-Buffering"), "nginx buffers the stream without it") +} + +func TestSSEBroker_TellsProxiesNotToBuffer(t *testing.T) { + broker := NewSSEBroker(slog.New(slog.NewTextHandler(io.Discard, nil))) + + ctx, cancel := context.WithCancel(context.Background()) + cancel() + rec := httptest.NewRecorder() + broker.ServeHTTP(rec, httptest.NewRequest(http.MethodGet, "/api/v1/containers/events", nil).WithContext(ctx)) + + assertUnbufferedEventStream(t, rec.Header()) +} + +type chunkLogRequester struct { + stubLogRequester + lines []string +} + +func (s *chunkLogRequester) SendCommand(context.Context, string, string, *agentpb.AgentCommand) (<-chan *agentpb.CommandResult, func(), error) { + results := make(chan *agentpb.CommandResult, 1) + results <- &agentpb.CommandResult{ + Last: true, + Result: &agentpb.CommandResult_Logs{Logs: &agentpb.LogsChunk{Lines: s.lines}}, + } + close(results) + return results, func() {}, nil +} + +func TestHandleLogStream_RemoteTellsProxiesNotToBuffer(t *testing.T) { + c := &container.Container{ + ID: "ctr-uuid", ExternalID: "deadbeef", Name: "web", + AgentID: "11111111-2222-3333-4444-555555555555", + } + svc := container.NewService(container.Deps{ + Store: &oneContainerStore{c: c}, + Logger: slog.New(slog.NewTextHandler(io.Discard, nil)), + }) + h := NewLogStreamHandler(nil, svc) + h.SetLogRequester(&chunkLogRequester{lines: []string{"hello"}}) + + req := httptest.NewRequest(http.MethodGet, "/api/v1/containers/ctr-uuid/logs/stream", nil) + req.SetPathValue("id", "ctr-uuid") + rec := httptest.NewRecorder() + h.HandleLogStream(rec, req) + + assert.Equal(t, http.StatusOK, rec.Code) + assertUnbufferedEventStream(t, rec.Header()) + assert.Contains(t, rec.Body.String(), "hello") +} diff --git a/internal/app/http.go b/internal/app/http.go index da272d86..ce67ec82 100644 --- a/internal/app/http.go +++ b/internal/app/http.go @@ -14,6 +14,7 @@ import ( "regexp" "strings" "time" + "unicode/utf8" "github.com/kolapsis/maintenant/cmd/maintenant/web" v1 "github.com/kolapsis/maintenant/internal/api/v1" @@ -187,17 +188,37 @@ func SecurityHeaders(csp string) func(http.Handler) http.Handler { } } +const minMCPClientSecretLength = 32 + +// warnWeakMCPClientSecret warns when the only credential the MCP OAuth flow checks is short enough to guess. +func (a *App) warnWeakMCPClientSecret() { + n := utf8.RuneCountInString(a.cfg.MCP.ClientSecret) + if n >= minMCPClientSecretLength { + return + } + a.logger.Warn("MAINTENANT_MCP_CLIENT_SECRET is shorter than 32 characters: /oauth/authorize approves every request, so this secret alone keeps /mcp closed", + "length", n, "minimum", minMCPClientSecretLength, "fix", "generate one with: openssl rand -hex 32") +} + +// unbufferedStream asks a buffering reverse proxy such as nginx to relay h's responses as they are written. +func unbufferedStream(h http.Handler) http.Handler { + return http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) { + w.Header().Set("X-Accel-Buffering", "no") + h.ServeHTTP(w, r) + }) +} + // buildHTTPServer assembles the top-level HTTP mux and creates the server. func (a *App) buildHTTPServer() *http.Server { topMux := http.NewServeMux() if a.cfg.MCP.Enabled && !a.cfg.DemoMode { - mcpHTTPHandler := gomcp.NewStreamableHTTPHandler(func(_ *http.Request) *gomcp.Server { + mcpHandler := unbufferedStream(gomcp.NewStreamableHTTPHandler(func(_ *http.Request) *gomcp.Server { return a.mcpServer - }, nil) - var mcpHandler http.Handler = mcpHTTPHandler + }, nil)) if a.cfg.MCP.ClientID != "" && a.cfg.MCP.ClientSecret != "" { + a.warnWeakMCPClientSecret() mcpOAuthStore := store.NewMCPOAuthStore(a.db) oauthSrv := mcpoauth.NewOAuthServer(mcpoauth.Config{ ClientID: a.cfg.MCP.ClientID, @@ -223,7 +244,7 @@ func (a *App) buildHTTPServer() *http.Server { authMiddleware := mcpauth.RequireBearerToken(tokenVerifier, &mcpauth.RequireBearerTokenOptions{ ResourceMetadataURL: resourceMetadataURL, }) - mcpHandler = authMiddleware(mcpHTTPHandler) + mcpHandler = authMiddleware(mcpHandler) go mcpoauth.StartCleanup(context.Background(), mcpOAuthStore, a.logger.With("component", "mcp-oauth-cleanup")) diff --git a/internal/app/http_wiring_test.go b/internal/app/http_wiring_test.go new file mode 100644 index 00000000..eae13969 --- /dev/null +++ b/internal/app/http_wiring_test.go @@ -0,0 +1,171 @@ +// Copyright 2026 Benjamin Touchard (Kolapsis) +// SPDX-License-Identifier: Apache-2.0 + +package app + +import ( + "context" + "encoding/json" + "io" + "net/http" + "net/http/httptest" + "net/url" + "strings" + "testing" + + "github.com/stretchr/testify/assert" + "github.com/stretchr/testify/require" + + v1 "github.com/kolapsis/maintenant/internal/api/v1" +) + +const wiringWebhookBody = `{"name":"hook","url":"https://8.8.8.8/hook","event_types":["*"]}` + +var crossSite = map[string]string{ + "Sec-Fetch-Site": "cross-site", + "Origin": "https://evil.example", + "Content-Type": "text/plain", +} + +func newTestApp(t *testing.T, edit func(*Config)) (*App, *syncBuffer) { + t.Helper() + cfg, logs, logger := storageEnv(t) + if edit != nil { + edit(&cfg) + } + a, err := New(cfg, logger) + require.NoError(t, err) + ctx, cancel := context.WithCancel(context.Background()) + a.db.StartWriter(ctx) + t.Cleanup(func() { + cancel() + _ = a.db.Close() + }) + return a, logs +} + +func serve(a *App, method, target, body string, headers map[string]string) *httptest.ResponseRecorder { + var r io.Reader + if body != "" { + r = strings.NewReader(body) + } + req := httptest.NewRequest(method, target, r) + for k, v := range headers { + req.Header.Set(k, v) + } + rec := httptest.NewRecorder() + a.srv.Handler.ServeHTTP(rec, req) + return rec +} + +func errorCode(t *testing.T, rec *httptest.ResponseRecorder) string { + t.Helper() + var body v1.ErrorResponse + if json.Unmarshal(rec.Body.Bytes(), &body) != nil { + return "" + } + return body.Error.Code +} + +func TestHTTPServer_CrossOriginGuardCoversTheAPIOnly(t *testing.T) { + a, _ := newTestApp(t, nil) + + rec := serve(a, http.MethodPost, "/api/v1/webhooks", wiringWebhookBody, crossSite) + assert.Equal(t, http.StatusForbidden, rec.Code) + assert.Equal(t, "CROSS_ORIGIN_REFUSED", errorCode(t, rec)) + + rec = serve(a, http.MethodPost, "/api/v1/webhooks", wiringWebhookBody, + map[string]string{"Sec-Fetch-Site": "same-origin", "Content-Type": "application/json"}) + assert.Equal(t, http.StatusCreated, rec.Code, rec.Body.String()) + + for _, target := range []string{"/ping/0b7c3a4e-7d1f-4c62-9a55-3f3b8d0c2e11", "/status/subscribe"} { + rec = serve(a, http.MethodPost, target, `{}`, crossSite) + assert.NotEqual(t, "CROSS_ORIGIN_REFUSED", errorCode(t, rec), "%s is called cross-origin by design", target) + } +} + +func TestHTTPServer_DemoDriverStillWrites(t *testing.T) { + a, _ := newTestApp(t, func(c *Config) { + c.DemoMode = true + c.DemoToken = "demo-driver-token" + }) + + for name, headers := range map[string]map[string]string{ + "seeding job": {v1.DemoTokenHeader: "demo-driver-token"}, + "same-origin page": {v1.DemoTokenHeader: "demo-driver-token", "Sec-Fetch-Site": "same-origin"}, + } { + t.Run(name, func(t *testing.T) { + rec := serve(a, http.MethodPost, "/api/v1/webhooks", wiringWebhookBody, headers) + assert.Equal(t, http.StatusCreated, rec.Code, rec.Body.String()) + }) + } + + rec := serve(a, http.MethodPost, "/api/v1/webhooks", wiringWebhookBody, nil) + assert.Equal(t, "DEMO_MODE", errorCode(t, rec)) +} + +const initializeRequest = `{"jsonrpc":"2.0","id":1,"method":"initialize","params":{"protocolVersion":"2025-06-18","capabilities":{},"clientInfo":{"name":"test","version":"1"}}}` + +func TestHTTPServer_MCPStreamIsNotBuffered(t *testing.T) { + a, _ := newTestApp(t, func(c *Config) { + c.MCP.Enabled = true + c.MCP.AllowUnauthenticated = true + }) + + rec := serve(a, http.MethodPost, "/mcp", initializeRequest, map[string]string{ + "Content-Type": "application/json", + "Accept": "application/json, text/event-stream", + }) + + require.Equal(t, http.StatusOK, rec.Code, rec.Body.String()) + assert.Equal(t, "text/event-stream", rec.Header().Get("Content-Type")) + assert.Equal(t, "no", rec.Header().Get("X-Accel-Buffering")) +} + +func TestHTTPServer_WarnsAboutAShortMCPClientSecret(t *testing.T) { + for name, tc := range map[string]struct { + secret string + warned bool + }{ + "short": {"hunter2", true}, + "31 chars": {strings.Repeat("a", 31), true}, + "32 chars": {strings.Repeat("a", 32), false}, + "openssl -hex": {strings.Repeat("0f", 32), false}, + } { + t.Run(name, func(t *testing.T) { + a, logs := newTestApp(t, func(c *Config) { + c.MCP.Enabled = true + c.MCP.ClientID = "claude" + c.MCP.ClientSecret = tc.secret + c.BaseURL = "https://maintenant.example.com" + }) + require.NotNil(t, a.srv, "a weak secret warns, it never stops the server") + + if tc.warned { + assert.Contains(t, logs.String(), "level=WARN") + assert.Contains(t, logs.String(), "MAINTENANT_MCP_CLIENT_SECRET is shorter than 32 characters") + } else { + assert.NotContains(t, logs.String(), "MAINTENANT_MCP_CLIENT_SECRET is shorter") + } + assert.NotContains(t, logs.String(), tc.secret, "the secret never reaches the logs") + }) + } +} + +func TestHTTPServer_OAuthEndpointsStayOpenCrossOrigin(t *testing.T) { + a, _ := newTestApp(t, func(c *Config) { + c.MCP.Enabled = true + c.MCP.ClientID = "claude" + c.MCP.ClientSecret = strings.Repeat("0f", 32) + c.BaseURL = "https://maintenant.example.com" + }) + + form := url.Values{"grant_type": {"authorization_code"}, "client_id": {"claude"}}.Encode() + rec := serve(a, http.MethodPost, "/oauth/token", form, map[string]string{ + "Sec-Fetch-Site": "cross-site", + "Origin": "https://claude.ai", + "Content-Type": "application/x-www-form-urlencoded", + }) + + assert.NotEqual(t, http.StatusForbidden, rec.Code, rec.Body.String()) +} diff --git a/internal/store/db_postgres.go b/internal/store/db_postgres.go index 487c3816..2c928093 100644 --- a/internal/store/db_postgres.go +++ b/internal/store/db_postgres.go @@ -6,9 +6,12 @@ package store import ( "context" "database/sql" + "encoding/binary" "errors" "fmt" + "io" "log/slog" + "net" "time" "github.com/jackc/pgx/v5/pgconn" @@ -46,13 +49,13 @@ func OpenPostgres(ctx context.Context, dsn string, logger *slog.Logger) (*DB, er defer cancel() if err := db.PingContext(pingCtx); err != nil { _ = db.Close() - return nil, classifyOpenError(err) + return nil, classifyOpenError(ctx, err) } var versionNum int if err := db.QueryRowContext(ctx, "SHOW server_version_num").Scan(&versionNum); err != nil { _ = db.Close() - return nil, fmt.Errorf("read server version: %w", classifyOpenError(err)) + return nil, fmt.Errorf("read server version: %w", classifyOpenError(ctx, err)) } if err := checkServerVersion(versionNum); err != nil { _ = db.Close() @@ -114,7 +117,7 @@ func checkServerVersion(versionNum int) error { // classifyOpenError maps a connection failure onto the startup sentinels so // the operator can tell refused credentials from an unreachable host. The // original error is kept in the chain; pgx never puts the password in it. -func classifyOpenError(err error) error { +func classifyOpenError(ctx context.Context, err error) error { var pe *pgconn.PgError if errors.As(err, &pe) { // 28P01 invalid_password, 28000 invalid_authorization_specification. @@ -122,9 +125,65 @@ func classifyOpenError(err error) error { return fmt.Errorf("%w: %s", ErrAuthRefused, pe.Message) } } + if tlsRefused(ctx, err) { + return fmt.Errorf("%w: %v", ErrTLSRefused, err) + } return fmt.Errorf("%w: %v", ErrUnreachable, err) } +const ( + sslRequestCode = 80877103 + tlsProbeTimeout = 3 * time.Second +) + +// tlsRefused reports whether err is a connection that had to use TLS failing on a server that declines TLS. +func tlsRefused(ctx context.Context, err error) bool { + var ce *pgconn.ConnectError + if !errors.As(err, &ce) || ce.Config == nil || !requiresTLS(ce.Config) { + return false + } + // pgx reports the refusal as a bare string, so the server is asked again. + return declinesTLS(ctx, ce.Config.Host, ce.Config.Port) +} + +func requiresTLS(cfg *pgconn.Config) bool { + if cfg.TLSConfig == nil { + return false + } + for _, fb := range cfg.Fallbacks { + if fb.TLSConfig == nil { + return false + } + } + return true +} + +// declinesTLS sends the protocol's SSLRequest, and nothing else, and reports whether the server answers 'N'. +func declinesTLS(ctx context.Context, host string, port uint16) bool { + ctx, cancel := context.WithTimeout(ctx, tlsProbeTimeout) + defer cancel() + + network, address := pgconn.NetworkAddress(host, port) + var d net.Dialer + conn, err := d.DialContext(ctx, network, address) + if err != nil { + return false + } + defer func() { _ = conn.Close() }() + if deadline, ok := ctx.Deadline(); ok { + _ = conn.SetDeadline(deadline) + } + + if err := binary.Write(conn, binary.BigEndian, []int32{8, sslRequestCode}); err != nil { + return false + } + var answer [1]byte + if _, err := io.ReadFull(conn, answer[:]); err != nil { + return false + } + return answer[0] == 'N' +} + // RedactedDSN exposes the connection target without its credentials, for the // startup log line and error messages. func (d *DB) RedactedDSN() string { diff --git a/internal/store/db_postgres_test.go b/internal/store/db_postgres_test.go index 173cc853..c077a0c6 100644 --- a/internal/store/db_postgres_test.go +++ b/internal/store/db_postgres_test.go @@ -5,6 +5,8 @@ package store import ( "context" + "encoding/binary" + "net" "net/url" "testing" "time" @@ -70,6 +72,65 @@ func TestOpenPostgres_AuthRefused(t *testing.T) { assert.NotContains(t, err.Error(), "definitely-not-the-password") } +// sslRequestServer answers each SSLRequest as PostgreSQL does with ssl=off ('N') or ssl=on ('S'), then hangs up. +func sslRequestServer(t *testing.T, answer byte) string { + t.Helper() + ln, err := net.Listen("tcp", "127.0.0.1:0") + require.NoError(t, err) + t.Cleanup(func() { _ = ln.Close() }) + + go func() { + for { + conn, err := ln.Accept() + if err != nil { + return + } + go func() { + defer func() { _ = conn.Close() }() + var req [2]int32 + if binary.Read(conn, binary.BigEndian, &req) != nil || req != [2]int32{8, sslRequestCode} { + return + } + _, _ = conn.Write([]byte{answer}) + }() + } + }() + return ln.Addr().String() +} + +func TestOpenPostgres_TLSRefused(t *testing.T) { + const password = "s3cr3t-Sentinel-tls" + addr := sslRequestServer(t, 'N') + + db, err := OpenPostgres(context.Background(), + "postgres://app:"+password+"@"+addr+"/maintenant?sslmode=require", testLogger()) + + require.ErrorIs(t, err, ErrTLSRefused) + assert.Nil(t, db) + assert.NotErrorIs(t, err, ErrUnreachable, "the server answered") + assert.NotContains(t, err.Error(), password) +} + +func TestOpenPostgres_TLSOfferedIsNotRefused(t *testing.T) { + addr := sslRequestServer(t, 'S') + + _, err := OpenPostgres(context.Background(), + "postgres://app:pw@"+addr+"/maintenant?sslmode=require", testLogger()) + + require.ErrorIs(t, err, ErrUnreachable, "a server offering TLS did not refuse it") + assert.NotErrorIs(t, err, ErrTLSRefused) +} + +func TestOpenPostgres_OptionalTLSIsNeverBlamed(t *testing.T) { + addr := sslRequestServer(t, 'N') + + _, err := OpenPostgres(context.Background(), + "postgres://app:pw@"+addr+"/maintenant?sslmode=prefer", testLogger()) + + require.ErrorIs(t, err, ErrUnreachable, "pgx falls back to plain text, so TLS is not the failure") + assert.NotErrorIs(t, err, ErrTLSRefused) +} + // TestCheckServerVersion covers the version refusal without needing an old // server to run against. func TestCheckServerVersion(t *testing.T) { diff --git a/internal/store/dsn.go b/internal/store/dsn.go index 95e8cd49..c8a1a427 100644 --- a/internal/store/dsn.go +++ b/internal/store/dsn.go @@ -50,19 +50,25 @@ func localDSNHost(u *url.URL) bool { // An unparseable string is returned unchanged; opening it fails with // ErrInvalidDSN anyway. func ApplyDefaultSSLMode(raw string) string { - u, err := ParseDSN(raw) - if err != nil { + if !DefaultsSSLMode(raw) { return raw } + u, _ := ParseDSN(raw) q := u.Query() - if q.Has("sslmode") || localDSNHost(u) { - return raw - } q.Set("sslmode", "require") u.RawQuery = q.Encode() return u.String() } +// DefaultsSSLMode reports whether ApplyDefaultSSLMode adds sslmode=require to raw. +func DefaultsSSLMode(raw string) bool { + u, err := ParseDSN(raw) + if err != nil { + return false + } + return !u.Query().Has("sslmode") && !localDSNHost(u) +} + // RedactDSN renders a connection string safe for logs and errors: // scheme://user@host:port/database, no password, no parameters. Principle VI // and FR-021: credentials never reach logs, responses or telemetry. diff --git a/internal/store/dsn_test.go b/internal/store/dsn_test.go index 41de9971..cf8b3122 100644 --- a/internal/store/dsn_test.go +++ b/internal/store/dsn_test.go @@ -62,6 +62,22 @@ func TestApplyDefaultSSLMode(t *testing.T) { assert.Equal(t, "not-a-dsn", ApplyDefaultSSLMode("not-a-dsn")) } +func TestDefaultsSSLMode_AgreesWithApplyDefaultSSLMode(t *testing.T) { + for _, raw := range []string{ + "postgres://app:pw@db.internal:5432/maintenant", + "postgres://app:pw@db:5432/maintenant", + "postgres://app@localhost:5432/maintenant", + "postgres://app@127.0.0.1/maintenant", + "postgres://app@/maintenant?host=/var/run/postgresql", + "postgres://app@db.internal/maintenant?sslmode=disable", + "postgres://app@db.internal/maintenant?sslmode=require", + "not-a-dsn", + } { + assert.Equal(t, ApplyDefaultSSLMode(raw) != raw, DefaultsSSLMode(raw), raw) + } + assert.True(t, DefaultsSSLMode("postgres://app:pw@db:5432/maintenant"), "a Compose service name is not local") +} + func TestRedactDSN(t *testing.T) { raw := "postgres://maintenant:" + secretPassword + "@db.internal:5432/prod?sslmode=require&application_name=x" red := RedactDSN(raw) diff --git a/internal/store/errors.go b/internal/store/errors.go index b6d721aa..06650339 100644 --- a/internal/store/errors.go +++ b/internal/store/errors.go @@ -26,6 +26,8 @@ var ( ErrUnreachable = errors.New("database unreachable") // ErrAuthRefused: the database answered and refused the credentials. ErrAuthRefused = errors.New("database credentials refused") + // ErrTLSRefused: the connection requires TLS and the server does not offer it. + ErrTLSRefused = errors.New("database server refused TLS") // ErrUnsupportedVersion: the server runs a version older than the minimum. ErrUnsupportedVersion = errors.New("database version unsupported (PostgreSQL 14 or newer required)") // ErrSchemaNewer: the schema was written by a newer release of this binary. diff --git a/internal/store/errors_test.go b/internal/store/errors_test.go index f245e1ea..146c4d93 100644 --- a/internal/store/errors_test.go +++ b/internal/store/errors_test.go @@ -80,7 +80,7 @@ func TestIsUnavailable(t *testing.T) { // TestSentinelsCarryNoSecret pins that no startup sentinel can ever leak a // credential: their messages are constants. func TestSentinelsCarryNoSecret(t *testing.T) { - for _, err := range []error{ErrInvalidDSN, ErrUnreachable, ErrAuthRefused, ErrUnsupportedVersion, ErrSchemaNewer} { + for _, err := range []error{ErrInvalidDSN, ErrUnreachable, ErrAuthRefused, ErrTLSRefused, ErrUnsupportedVersion, ErrSchemaNewer} { assert.NotContains(t, err.Error(), "://") assert.NotContains(t, err.Error(), "password") } From 924eefd25d8c8a8f7891cb1cc845b6b1b251b82a Mon Sep 17 00:00:00 2001 From: Benjamin Date: Wed, 30 Sep 2026 16:25:35 +0200 Subject: [PATCH 04/54] fix(deploy): make the image, entrypoint and Kubernetes manifests work as shipped - Entrypoint: exec the command directly when not root (runAsUser, --user) instead of failing in setpriv; detect agent mode from MAINTENANT_MODE and the real data dir; own the database directory, never its content nor /. - Image: default MAINTENANT_DB to /data/maintenant.db, on the data volume. - Config: read MAINTENANT_NODE_NAME from the environment. - Kubernetes runtime: Connect retries until the API server answers, and the local topology reconcile only runs once connected, which removes the nil pointer panic on an unreachable cluster. - RBAC: one read-only rule list derived from the runtime's API calls, used by the agent manifest and checked against the server manifests (adds nodes, batch/jobs, metrics nodes; drops replicasets). - Agent manifest: a single-replica Deployment with a PVC for identity and spool, hardened, instead of a DaemonSet on hostPath. - deploy/kubernetes: pin the maintenant namespace on namespaced objects. - Standalone install tab: hand out the install.maintenant.dev command. --- Dockerfile | 3 + cmd/maintenant/image_test.go | 153 ++++++++++++ deploy/helm/maintenant/templates/rbac.yaml | 17 +- deploy/kubernetes/deployment.yaml | 9 +- deploy/kubernetes/rbac.yaml | 29 ++- docker-entrypoint.sh | 45 +++- .../agents/EnrollmentTokenModal.vue | 3 +- .../__tests__/EnrollmentTokenModal.spec.ts | 49 ++++ internal/api/v1/agents_handler.go | 77 ++++-- internal/api/v1/agents_token_handler_test.go | 7 +- internal/api/v1/install_templates_test.go | 174 +++++++++++--- internal/app/app.go | 6 +- internal/app/config.go | 1 + internal/app/config_test.go | 6 + internal/app/lifecycle.go | 16 +- internal/app/lifecycle_test.go | 35 +++ internal/kubernetes/rbac.go | 18 ++ internal/kubernetes/rbac_test.go | 219 ++++++++++++++++++ internal/kubernetes/runtime.go | 35 ++- internal/kubernetes/runtime_test.go | 90 +++++++ scripts/ci-changed-paths.sh | 10 +- 21 files changed, 890 insertions(+), 112 deletions(-) create mode 100644 cmd/maintenant/image_test.go create mode 100644 frontend/src/commercial/components/agents/__tests__/EnrollmentTokenModal.spec.ts create mode 100644 internal/kubernetes/rbac.go create mode 100644 internal/kubernetes/rbac_test.go create mode 100644 internal/kubernetes/runtime_test.go diff --git a/Dockerfile b/Dockerfile index ba75e9de..37f0a3f9 100644 --- a/Dockerfile +++ b/Dockerfile @@ -51,6 +51,9 @@ RUN apk add --no-cache ca-certificates tzdata setpriv \ # /tmp as a tiny tmpfs, which SQLITE_FULL-fails the conversion; /data has real space. ENV SQLITE_TMPDIR=/data +# Its directory also holds the licence cache and the update window, PostgreSQL or not. +ENV MAINTENANT_DB=/data/maintenant.db + # Tells the OS identity reader it must not fall back to the image's own # /etc/os-release, which describes the container rather than the host. ENV MAINTENANT_CONTAINER=1 diff --git a/cmd/maintenant/image_test.go b/cmd/maintenant/image_test.go new file mode 100644 index 00000000..e8ed7a3b --- /dev/null +++ b/cmd/maintenant/image_test.go @@ -0,0 +1,153 @@ +// Copyright 2026 Benjamin Touchard (kOlapsis) +// SPDX-License-Identifier: Apache-2.0 + +package main + +import ( + "os" + "os/exec" + "path/filepath" + "strings" + "testing" + + "github.com/stretchr/testify/assert" + "github.com/stretchr/testify/require" +) + +func TestImageKeepsTheDatabaseOnTheDataVolume(t *testing.T) { + dockerfile, err := os.ReadFile(filepath.Join("..", "..", "Dockerfile")) + require.NoError(t, err) + + env := map[string]string{} + var volumes []string + for _, line := range strings.Split(string(dockerfile), "\n") { + fields := strings.Fields(line) + if len(fields) < 2 { + continue + } + switch fields[0] { + case "ENV": + for _, kv := range fields[1:] { + if k, v, ok := strings.Cut(kv, "="); ok { + env[k] = v + } + } + case "VOLUME": + volumes = append(volumes, fields[1:]...) + } + } + + db := env["MAINTENANT_DB"] + require.NotEmpty(t, db, "without MAINTENANT_DB the binary opens ./maintenant.db on the read-only root") + assert.Contains(t, volumes, filepath.Dir(db)) +} + +type entrypointRun struct { + stubs string + log string +} + +// newEntrypointRun stubs every command the entrypoint runs, so it can be driven as root without touching the host. +func newEntrypointRun(t *testing.T) *entrypointRun { + t.Helper() + stubs := t.TempDir() + record := "#!/bin/sh\nprintf '%s %s\\n' \"${0##*/}\" \"$*\" >>\"$STUB_LOG\"\n" + for _, name := range []string{"chown", "mkdir", "setpriv", "stat", "maintenant"} { + require.NoError(t, os.WriteFile(filepath.Join(stubs, name), []byte(record), 0o700)) + } + id := "#!/bin/sh\necho \"$STUB_UID\"\n" + require.NoError(t, os.WriteFile(filepath.Join(stubs, "id"), []byte(id), 0o700)) + return &entrypointRun{stubs: stubs, log: filepath.Join(stubs, "calls.log")} +} + +func (e *entrypointRun) binary() string { return filepath.Join(e.stubs, "maintenant") } + +// run executes the entrypoint as uid and returns the stub calls, one "name args" per entry. +func (e *entrypointRun) run(t *testing.T, uid string, env map[string]string, args ...string) []string { + t.Helper() + cmd := exec.Command("/bin/sh", append([]string{filepath.Join("..", "..", "docker-entrypoint.sh")}, args...)...) + cmd.Env = []string{"PATH=" + e.stubs + ":" + os.Getenv("PATH"), "STUB_LOG=" + e.log, "STUB_UID=" + uid} + for k, v := range env { + cmd.Env = append(cmd.Env, k+"="+v) + } + out, err := cmd.CombinedOutput() + require.NoError(t, err, string(out)) + data, err := os.ReadFile(e.log) + require.NoError(t, err) + return strings.Split(strings.TrimSpace(string(data)), "\n") +} + +func calls(all []string, name string) []string { + var out []string + for _, c := range all { + if cmd, args, _ := strings.Cut(c, " "); cmd == name { + out = append(out, args) + } + } + return out +} + +func chowned(all []string) []string { + var dirs []string + for _, args := range calls(all, "chown") { + f := strings.Fields(args) + dirs = append(dirs, f[len(f)-1]) + } + return dirs +} + +func TestEntrypoint_UnprivilegedRunsTheCommandItself(t *testing.T) { + e := newEntrypointRun(t) + got := e.run(t, "65534", nil, e.binary(), "--mode=agent") + + assert.Equal(t, []string{"maintenant --mode=agent"}, got) +} + +func TestEntrypoint_AgentOwnsItsDataDirectory(t *testing.T) { + dataDir := t.TempDir() + cases := map[string]struct { + env map[string]string + args []string + }{ + "mode and data dir from the environment": { + env: map[string]string{"MAINTENANT_MODE": "agent", "MAINTENANT_DATA_DIR": dataDir}, + }, + "mode and data dir from flags": { + args: []string{"--mode", "agent", "--data-dir=" + dataDir}, + }, + "mode flag, data dir from the environment": { + env: map[string]string{"MAINTENANT_DATA_DIR": dataDir}, + args: []string{"-mode=agent"}, + }, + } + for name, tc := range cases { + t.Run(name, func(t *testing.T) { + e := newEntrypointRun(t) + got := e.run(t, "0", tc.env, append([]string{e.binary()}, tc.args...)...) + + assert.Equal(t, []string{dataDir}, chowned(got)) + assert.Len(t, calls(got, "setpriv"), 1) + }) + } +} + +func TestEntrypoint_ServerOwnsTheDatabaseDirectory(t *testing.T) { + dbDir := t.TempDir() + flagDir := t.TempDir() + + e := newEntrypointRun(t) + got := e.run(t, "0", map[string]string{"MAINTENANT_DB": filepath.Join(dbDir, "maintenant.db")}, e.binary()) + assert.ElementsMatch(t, []string{dbDir, "/data/shm"}, chowned(got)) + assert.Len(t, calls(got, "setpriv"), 1) + + e = newEntrypointRun(t) + got = e.run(t, "0", map[string]string{"MAINTENANT_MODE": "agent"}, e.binary(), "--mode=server", "--db", filepath.Join(flagDir, "m.db")) + assert.ElementsMatch(t, []string{flagDir, "/data/shm"}, chowned(got)) +} + +func TestEntrypoint_NeverHandsTheRootDirectoryOver(t *testing.T) { + e := newEntrypointRun(t) + got := e.run(t, "0", map[string]string{"MAINTENANT_DB": "/maintenant.db"}, e.binary()) + + assert.Equal(t, []string{"/data/shm"}, chowned(got)) +} diff --git a/deploy/helm/maintenant/templates/rbac.yaml b/deploy/helm/maintenant/templates/rbac.yaml index 8fe6aec0..6994a5bd 100644 --- a/deploy/helm/maintenant/templates/rbac.yaml +++ b/deploy/helm/maintenant/templates/rbac.yaml @@ -6,17 +6,22 @@ metadata: labels: {{- include "maintenant.labels" . | nindent 4 }} rules: - # Core resources — read-only + # Read-only, exactly the APIs the runtime calls (internal/kubernetes/rbac.go). - apiGroups: [""] - resources: ["pods", "pods/log", "services", "namespaces", "events"] + resources: ["namespaces", "nodes", "pods", "events", "services"] verbs: ["get", "list", "watch"] - # Apps — read-only + - apiGroups: [""] + resources: ["pods/log"] + verbs: ["get"] - apiGroups: ["apps"] - resources: ["deployments", "statefulsets", "daemonsets", "replicasets"] + resources: ["deployments", "statefulsets", "daemonsets"] + verbs: ["get", "list", "watch"] + - apiGroups: ["batch"] + resources: ["jobs"] verbs: ["get", "list", "watch"] - # Metrics — read-only (requires metrics-server) + # Requires metrics-server. - apiGroups: ["metrics.k8s.io"] - resources: ["pods"] + resources: ["pods", "nodes"] verbs: ["get", "list"] --- apiVersion: rbac.authorization.k8s.io/v1 diff --git a/deploy/kubernetes/deployment.yaml b/deploy/kubernetes/deployment.yaml index c607556b..426f23a9 100644 --- a/deploy/kubernetes/deployment.yaml +++ b/deploy/kubernetes/deployment.yaml @@ -1,9 +1,12 @@ -# Apply with: kubectl apply -n -f deployment.yaml -# No namespace is hardcoded — kubectl -n sets it. +# Apply with: +# kubectl create namespace maintenant +# kubectl apply -f deploy/kubernetes/ +# The namespace is fixed here and in rbac.yaml; change both together. apiVersion: apps/v1 kind: Deployment metadata: name: maintenant + namespace: maintenant labels: app.kubernetes.io/name: maintenant app.kubernetes.io/component: monitoring @@ -105,6 +108,7 @@ apiVersion: v1 kind: PersistentVolumeClaim metadata: name: maintenant-data + namespace: maintenant spec: accessModes: - ReadWriteOnce @@ -119,6 +123,7 @@ apiVersion: v1 kind: Service metadata: name: maintenant + namespace: maintenant labels: app.kubernetes.io/name: maintenant spec: diff --git a/deploy/kubernetes/rbac.yaml b/deploy/kubernetes/rbac.yaml index 7634d2be..7213794f 100644 --- a/deploy/kubernetes/rbac.yaml +++ b/deploy/kubernetes/rbac.yaml @@ -1,28 +1,35 @@ -# Apply with: kubectl apply -n -f rbac.yaml -# The ServiceAccount is created in the target namespace automatically. -# The ClusterRoleBinding references it, so update the namespace below -# if you are NOT using "maintenant" as your namespace. +# Apply with: +# kubectl create namespace maintenant +# kubectl apply -f deploy/kubernetes/ +# The namespace is fixed here, in the ClusterRoleBinding subject and in +# deployment.yaml; change all of them together. apiVersion: v1 kind: ServiceAccount metadata: name: maintenant + namespace: maintenant --- apiVersion: rbac.authorization.k8s.io/v1 kind: ClusterRole metadata: name: maintenant rules: - # Core resources — read-only + # Read-only, exactly the APIs the runtime calls (internal/kubernetes/rbac.go). - apiGroups: [""] - resources: ["pods", "pods/log", "services", "namespaces", "events"] + resources: ["namespaces", "nodes", "pods", "events", "services"] verbs: ["get", "list", "watch"] - # Apps — read-only + - apiGroups: [""] + resources: ["pods/log"] + verbs: ["get"] - apiGroups: ["apps"] - resources: ["deployments", "statefulsets", "daemonsets", "replicasets"] + resources: ["deployments", "statefulsets", "daemonsets"] + verbs: ["get", "list", "watch"] + - apiGroups: ["batch"] + resources: ["jobs"] verbs: ["get", "list", "watch"] - # Metrics — read-only (requires metrics-server) + # Requires metrics-server. - apiGroups: ["metrics.k8s.io"] - resources: ["pods"] + resources: ["pods", "nodes"] verbs: ["get", "list"] --- apiVersion: rbac.authorization.k8s.io/v1 @@ -36,4 +43,4 @@ roleRef: subjects: - kind: ServiceAccount name: maintenant - namespace: maintenant # ← change this to match your namespace + namespace: maintenant diff --git a/docker-entrypoint.sh b/docker-entrypoint.sh index 2f3ebf14..a8feffe1 100755 --- a/docker-entrypoint.sh +++ b/docker-entrypoint.sh @@ -8,18 +8,43 @@ if [ "${1#-}" != "$1" ]; then set -- /app/maintenant "$@" fi -# Ensure the data directory used by the chosen mode is writable by the -# unprivileged runtime user. Covers both named volumes (Docker creates them -# root:root when the target path doesn't exist in the image) and bind mounts -# (host ownership leaks into the container). -case " $* " in - *" --mode=agent "*|*" --mode agent "*) - mkdir -p /var/lib/maintenant - chown 65534:65534 /var/lib/maintenant +# Already unprivileged (runAsUser, --user): no chown possible, and setpriv needs CAP_SETGID. +if [ "$(id -u)" != "0" ]; then + exec "$@" +fi + +mode="${MAINTENANT_MODE:-embedded}" +db="${MAINTENANT_DB:-./maintenant.db}" +data_dir="${MAINTENANT_DATA_DIR:-/var/lib/maintenant}" +prev="" +for arg in "$@"; do + case "$prev" in + --mode | -mode) mode="$arg" ;; + --db | -db) db="$arg" ;; + --data-dir | -data-dir) data_dir="$arg" ;; + esac + case "$arg" in + --mode=* | -mode=*) mode="${arg#*=}" ;; + --db=* | -db=*) db="${arg#*=}" ;; + --data-dir=* | -data-dir=*) data_dir="${arg#*=}" ;; + esac + prev="$arg" +done + +# Volumes and bind mounts arrive root-owned: hand the directory itself, never its content nor /, to the runtime user. +own_dir() { + mkdir -p -- "$1" + [ "$(cd -- "$1" && pwd -P)" != "/" ] || return 0 + chown 65534:65534 -- "$1" +} + +case "$mode" in + agent) + own_dir "$data_dir" ;; *) - mkdir -p /data/shm - chown 65534:65534 /data/shm + own_dir "$(dirname -- "$db")" + own_dir /data/shm ;; esac diff --git a/frontend/src/commercial/components/agents/EnrollmentTokenModal.vue b/frontend/src/commercial/components/agents/EnrollmentTokenModal.vue index bcf920d5..05e7a835 100644 --- a/frontend/src/commercial/components/agents/EnrollmentTokenModal.vue +++ b/frontend/src/commercial/components/agents/EnrollmentTokenModal.vue @@ -25,7 +25,7 @@ const MODES: Array<{ id: InstallMode; label: string }> = [ { id: 'docker_run', label: 'Docker run' }, { id: 'docker_compose', label: 'Compose' }, { id: 'kubernetes', label: 'Kubernetes' }, - { id: 'standalone', label: 'Standalone (soon)' }, + { id: 'standalone', label: 'Standalone' }, ] const MODE_OPTIONS = MODES.map((m) => ({ value: m.id, label: m.label })) @@ -126,7 +126,6 @@ watch(open, (value) => { >{{ currentTemplate }} { + wrapper?.unmount() + wrapper = undefined + window.localStorage.clear() +}) + +function buttonLabelled(label: string): HTMLButtonElement | undefined { + return Array.from(document.querySelectorAll('button')).find( + (b) => b.textContent?.trim() === label, + ) +} + +describe('EnrollmentTokenModal', () => { + it('hands out the standalone installer as a command to copy', async () => { + wrapper = mount(EnrollmentTokenModal, { props: { token }, attachTo: document.body }) + + const standalone = buttonLabelled('Standalone') + expect(standalone).toBeDefined() + standalone?.click() + await nextTick() + + expect(document.querySelector('pre')?.textContent).toBe(token.install_templates.standalone) + expect(buttonLabelled('Copy install command')).toBeDefined() + }) +}) diff --git a/internal/api/v1/agents_handler.go b/internal/api/v1/agents_handler.go index 59f5105e..cce90612 100644 --- a/internal/api/v1/agents_handler.go +++ b/internal/api/v1/agents_handler.go @@ -10,6 +10,7 @@ import ( "net/http" "net/netip" "strconv" + "strings" "time" "github.com/kolapsis/maintenant/internal/agent" @@ -17,6 +18,7 @@ import ( "github.com/kolapsis/maintenant/internal/eol" "github.com/kolapsis/maintenant/internal/event" "github.com/kolapsis/maintenant/internal/extension" + "github.com/kolapsis/maintenant/internal/kubernetes" "github.com/kolapsis/maintenant/internal/store" ) @@ -145,20 +147,18 @@ func (h *AgentHandler) HandleCreateEnrollmentToken(w http.ResponseWriter, r *htt func buildInstallTemplates(serverURL, token string) map[string]string { return map[string]string{ - "standalone": buildInstallStandalone(), + "standalone": buildInstallStandalone(serverURL, token), "docker_run": buildInstallDockerRun(serverURL, token), "docker_compose": buildInstallDockerCompose(serverURL, token), "kubernetes": buildInstallKubernetes(serverURL, token), } } -// The standalone installer is not released yet. Until it is, the tab announces -// itself rather than handing out an invocation that would fetch nothing. -func buildInstallStandalone() string { - return "Coming soon.\n\n" + - "The standalone installer (binary + systemd unit) is not released yet.\n" + - "Run the agent with Docker in the meantime — see the Docker run,\n" + - "Compose and Kubernetes tabs." +func buildInstallStandalone(serverURL, token string) string { + return "curl -fsSL https://install.maintenant.dev | sudo bash -s -- \\\n" + + " --mode agent \\\n" + + " --server " + serverURL + " \\\n" + + " --enrollment-token " + token } func buildInstallDockerRun(serverURL, token string) string { @@ -219,9 +219,7 @@ func buildInstallKubernetes(serverURL, token string) string { "metadata:\n" + " name: maintenant-agent\n" + "rules:\n" + - " - apiGroups: [\"\"]\n" + - " resources: [pods, nodes, services, events]\n" + - " verbs: [get, list, watch]\n" + + kubernetesReadRules() + "---\n" + "apiVersion: rbac.authorization.k8s.io/v1\n" + "kind: ClusterRoleBinding\n" + @@ -236,12 +234,26 @@ func buildInstallKubernetes(serverURL, token string) string { " name: maintenant-agent\n" + " namespace: maintenant\n" + "---\n" + + "apiVersion: v1\n" + + "kind: PersistentVolumeClaim\n" + + "metadata:\n" + + " name: maintenant-agent-data\n" + + " namespace: maintenant\n" + + "spec:\n" + + " accessModes: [ReadWriteOnce]\n" + + " resources:\n" + + " requests:\n" + + " storage: 1Gi\n" + + "---\n" + "apiVersion: apps/v1\n" + - "kind: DaemonSet\n" + + "kind: Deployment\n" + "metadata:\n" + " name: maintenant-agent\n" + " namespace: maintenant\n" + "spec:\n" + + " replicas: 1\n" + + " strategy:\n" + + " type: Recreate\n" + " selector:\n" + " matchLabels: { app: maintenant-agent }\n" + " template:\n" + @@ -249,13 +261,18 @@ func buildInstallKubernetes(serverURL, token string) string { " labels: { app: maintenant-agent }\n" + " spec:\n" + " serviceAccountName: maintenant-agent\n" + + " hostname: maintenant-agent\n" + + " securityContext:\n" + + " runAsNonRoot: true\n" + + " runAsUser: 65534\n" + + " runAsGroup: 65534\n" + + " fsGroup: 65534\n" + " containers:\n" + " - name: agent\n" + " image: ghcr.io/kolapsis/maintenant:latest\n" + " args:\n" + " - --mode=agent\n" + " - --server=" + serverURL + "\n" + - " - --enrollment-token=$(MAINTENANT_ENROLLMENT_TOKEN)\n" + " - --runtime=kubernetes\n" + " env:\n" + " - name: MAINTENANT_ENROLLMENT_TOKEN\n" + @@ -263,19 +280,37 @@ func buildInstallKubernetes(serverURL, token string) string { " secretKeyRef:\n" + " name: maintenant-agent-enrollment\n" + " key: token\n" + - " - name: MAINTENANT_LABEL\n" + - " valueFrom:\n" + - " fieldRef: { fieldPath: spec.nodeName }\n" + " - name: MAINTENANT_NODE_NAME\n" + " valueFrom:\n" + " fieldRef: { fieldPath: spec.nodeName }\n" + + " securityContext:\n" + + " allowPrivilegeEscalation: false\n" + + " readOnlyRootFilesystem: true\n" + + " capabilities:\n" + + " drop: [ALL]\n" + " volumeMounts:\n" + - " - { name: identity, mountPath: /var/lib/maintenant }\n" + + " - { name: data, mountPath: /var/lib/maintenant }\n" + + " - { name: tmp, mountPath: /tmp }\n" + " volumes:\n" + - " - name: identity\n" + - " hostPath:\n" + - " path: /var/lib/maintenant-agent\n" + - " type: DirectoryOrCreate\n" + " - name: data\n" + + " persistentVolumeClaim:\n" + + " claimName: maintenant-agent-data\n" + + " - name: tmp\n" + + " emptyDir: {}\n" +} + +func kubernetesReadRules() string { + var b strings.Builder + for _, r := range kubernetes.ReadRules() { + groups := make([]string, len(r.APIGroups)) + for i, g := range r.APIGroups { + groups[i] = strconv.Quote(g) + } + b.WriteString(" - apiGroups: [" + strings.Join(groups, ", ") + "]\n" + + " resources: [" + strings.Join(r.Resources, ", ") + "]\n" + + " verbs: [" + strings.Join(r.Verbs, ", ") + "]\n") + } + return b.String() } func (h *AgentHandler) HandleListEnrollmentTokens(w http.ResponseWriter, r *http.Request) { diff --git a/internal/api/v1/agents_token_handler_test.go b/internal/api/v1/agents_token_handler_test.go index 018a20fe..0466e3ab 100644 --- a/internal/api/v1/agents_token_handler_test.go +++ b/internal/api/v1/agents_token_handler_test.go @@ -64,16 +64,11 @@ func TestCreateEnrollmentToken_ReturnsCleartextOnce(t *testing.T) { assert.Equal(t, agent.TokenIDFromHash(agent.HashToken(cleartext)), out["token_id"]) // The install snippets have to embed the real token or the operator cannot - // copy-paste them. The standalone tab is the exception: it announces an - // unreleased installer instead of handing out a command. + // copy-paste them. templates, ok := out["install_templates"].(map[string]any) require.True(t, ok) require.NotEmpty(t, templates) for name, tmpl := range templates { - if name == "standalone" { - assert.NotContains(t, tmpl, cleartext, "the coming-soon notice must not carry the token") - continue - } assert.Contains(t, tmpl, cleartext, "install template %q must carry the token", name) } diff --git a/internal/api/v1/install_templates_test.go b/internal/api/v1/install_templates_test.go index 52b8f18c..6192c05a 100644 --- a/internal/api/v1/install_templates_test.go +++ b/internal/api/v1/install_templates_test.go @@ -4,8 +4,21 @@ package v1 import ( + "bytes" + "encoding/json" + "errors" + "io" "strings" "testing" + + "github.com/stretchr/testify/assert" + "github.com/stretchr/testify/require" + appsv1 "k8s.io/api/apps/v1" + corev1 "k8s.io/api/core/v1" + rbacv1 "k8s.io/api/rbac/v1" + "k8s.io/apimachinery/pkg/util/yaml" + + "github.com/kolapsis/maintenant/internal/kubernetes" ) func TestBuildInstallTemplates_ReturnsFourModes(t *testing.T) { @@ -23,14 +36,6 @@ func TestBuildInstallTemplates_ReturnsFourModes(t *testing.T) { t.Errorf("missing template key %q", k) continue } - if v == "" { - t.Errorf("template %q is empty", k) - } - // The standalone installer is unreleased: its tab carries an - // announcement, not a command, so it holds neither URL nor token. - if k == "standalone" { - continue - } if !strings.Contains(v, serverURL) { t.Errorf("template %q does not contain server URL %q", k, serverURL) } @@ -44,16 +49,13 @@ func TestBuildInstallTemplates_ReturnsFourModes(t *testing.T) { } } -func TestBuildInstallStandalone_AnnouncesComingSoon(t *testing.T) { - out := buildInstallStandalone() - if !strings.HasPrefix(out, "Coming soon.") { - t.Errorf("standalone template should announce the unreleased installer, got:\n%s", out) - } - // install.maintenant.dev serves the marketing site, so an invocation - // piping it into a shell would run HTML as root. - if strings.Contains(out, "install.maintenant.dev") { - t.Errorf("standalone template still points at the unpublished installer:\n%s", out) - } +func TestBuildInstallStandalone_PipesTheInstallerInAgentMode(t *testing.T) { + out := buildInstallStandalone("grpcs://h:8443", "mnt_enr_xyz") + + assert.True(t, strings.HasPrefix(out, "curl -fsSL https://install.maintenant.dev | sudo bash -s -- "), out) + assert.Contains(t, out, "--mode agent") + assert.Contains(t, out, "--server grpcs://h:8443") + assert.Contains(t, out, "--enrollment-token mnt_enr_xyz") } func TestBuildInstallDockerRun_ContainsDockerSocketMount(t *testing.T) { @@ -88,24 +90,124 @@ func TestBuildInstallDockerCompose_HasServicesBlock(t *testing.T) { } } -func TestBuildInstallKubernetes_HasDaemonSetAndSecret(t *testing.T) { - out := buildInstallKubernetes("grpcs://h:8443", "mnt_enr_xyz") - if !strings.Contains(out, "kind: DaemonSet") { - t.Errorf("kubernetes template missing DaemonSet kind") - } - if !strings.Contains(out, "kind: Secret") { - t.Errorf("kubernetes template missing Secret kind") - } - if !strings.Contains(out, "kind: ClusterRole") { - t.Errorf("kubernetes template missing ClusterRole (RBAC)") - } - if !strings.Contains(out, "maintenant-agent-enrollment") { - t.Errorf("kubernetes template missing enrollment secret name") - } - if !strings.Contains(out, "--runtime=kubernetes") { - t.Errorf("kubernetes template should pin runtime to kubernetes") +// kubernetesManifest decodes the generated multi-document manifest into its objects, by kind. +func kubernetesManifest(t *testing.T, out string) map[string][]json.RawMessage { + t.Helper() + dec := yaml.NewYAMLOrJSONDecoder(bytes.NewReader([]byte(out)), 4096) + objs := map[string][]json.RawMessage{} + for { + var raw json.RawMessage + err := dec.Decode(&raw) + if errors.Is(err, io.EOF) { + return objs + } + require.NoError(t, err) + var head struct { + Kind string `json:"kind"` + } + require.NoError(t, json.Unmarshal(raw, &head)) + objs[head.Kind] = append(objs[head.Kind], raw) } - if !strings.Contains(out, "name: MAINTENANT_NODE_NAME") { - t.Errorf("kubernetes template missing MAINTENANT_NODE_NAME env entry") +} + +func decodeOne[T any](t *testing.T, objs map[string][]json.RawMessage, kind string) T { + t.Helper() + var obj T + require.Len(t, objs[kind], 1, "want exactly one %s", kind) + dec := json.NewDecoder(bytes.NewReader(objs[kind][0])) + dec.DisallowUnknownFields() + require.NoError(t, dec.Decode(&obj), "%s has a field Kubernetes does not know", kind) + return obj +} + +func TestBuildInstallKubernetes_OneAgentPerCluster(t *testing.T) { + objs := kubernetesManifest(t, buildInstallKubernetes("grpcs://h:8443", "mnt_enr_xyz")) + + assert.Empty(t, objs["DaemonSet"], "one pod per node duplicates the cluster topology") + secret := decodeOne[corev1.Secret](t, objs, "Secret") + assert.Equal(t, "mnt_enr_xyz", secret.StringData["token"]) + + dep := decodeOne[appsv1.Deployment](t, objs, "Deployment") + require.NotNil(t, dep.Spec.Replicas) + assert.Equal(t, int32(1), *dep.Spec.Replicas) + assert.Equal(t, appsv1.RecreateDeploymentStrategyType, dep.Spec.Strategy.Type) + + pod := dep.Spec.Template.Spec + require.Len(t, pod.Containers, 1) + c := pod.Containers[0] + assert.Contains(t, c.Args, "--mode=agent") + assert.Contains(t, c.Args, "--server=grpcs://h:8443") + assert.Contains(t, c.Args, "--runtime=kubernetes") + + env := map[string]corev1.EnvVar{} + for _, e := range c.Env { + env[e.Name] = e + } + require.NotNil(t, env["MAINTENANT_ENROLLMENT_TOKEN"].ValueFrom) + assert.Equal(t, "maintenant-agent-enrollment", env["MAINTENANT_ENROLLMENT_TOKEN"].ValueFrom.SecretKeyRef.Name) + require.NotNil(t, env["MAINTENANT_NODE_NAME"].ValueFrom) + assert.Equal(t, "spec.nodeName", env["MAINTENANT_NODE_NAME"].ValueFrom.FieldRef.FieldPath) +} + +func TestBuildInstallKubernetes_StateSurvivesAReschedule(t *testing.T) { + objs := kubernetesManifest(t, buildInstallKubernetes("grpcs://h:8443", "mnt_enr_xyz")) + pvc := decodeOne[corev1.PersistentVolumeClaim](t, objs, "PersistentVolumeClaim") + dep := decodeOne[appsv1.Deployment](t, objs, "Deployment") + + assert.Equal(t, []corev1.PersistentVolumeAccessMode{corev1.ReadWriteOnce}, pvc.Spec.AccessModes) + assert.Equal(t, pvc.Namespace, dep.Namespace) + + pod := dep.Spec.Template.Spec + volumes := map[string]corev1.Volume{} + for _, v := range pod.Volumes { + assert.Nil(t, v.HostPath, "volume %q pins the agent state to one node", v.Name) + volumes[v.Name] = v + } + var dataMount *corev1.VolumeMount + for i, m := range pod.Containers[0].VolumeMounts { + if m.MountPath == "/var/lib/maintenant" { + dataMount = &pod.Containers[0].VolumeMounts[i] + } } + require.NotNil(t, dataMount, "nothing mounted on the agent data directory") + claim := volumes[dataMount.Name].PersistentVolumeClaim + require.NotNil(t, claim, "the agent data directory is not backed by a claim") + assert.Equal(t, pvc.Name, claim.ClaimName) +} + +func TestBuildInstallKubernetes_Hardened(t *testing.T) { + objs := kubernetesManifest(t, buildInstallKubernetes("grpcs://h:8443", "mnt_enr_xyz")) + dep := decodeOne[appsv1.Deployment](t, objs, "Deployment") + + podSC := dep.Spec.Template.Spec.SecurityContext + require.NotNil(t, podSC) + require.NotNil(t, podSC.RunAsNonRoot) + assert.True(t, *podSC.RunAsNonRoot) + require.NotNil(t, podSC.RunAsUser) + assert.Equal(t, int64(65534), *podSC.RunAsUser) + require.NotNil(t, podSC.FSGroup) + assert.Equal(t, int64(65534), *podSC.FSGroup) + + sc := dep.Spec.Template.Spec.Containers[0].SecurityContext + require.NotNil(t, sc) + require.NotNil(t, sc.ReadOnlyRootFilesystem) + assert.True(t, *sc.ReadOnlyRootFilesystem) + require.NotNil(t, sc.AllowPrivilegeEscalation) + assert.False(t, *sc.AllowPrivilegeEscalation) + require.NotNil(t, sc.Capabilities) + assert.Equal(t, []corev1.Capability{"ALL"}, sc.Capabilities.Drop) +} + +func TestBuildInstallKubernetes_GrantsWhatTheRuntimeReads(t *testing.T) { + objs := kubernetesManifest(t, buildInstallKubernetes("grpcs://h:8443", "mnt_enr_xyz")) + + role := decodeOne[rbacv1.ClusterRole](t, objs, "ClusterRole") + assert.Equal(t, kubernetes.ReadRules(), role.Rules) + + sa := decodeOne[corev1.ServiceAccount](t, objs, "ServiceAccount") + binding := decodeOne[rbacv1.ClusterRoleBinding](t, objs, "ClusterRoleBinding") + assert.Equal(t, role.Name, binding.RoleRef.Name) + require.Len(t, binding.Subjects, 1) + assert.Equal(t, sa.Name, binding.Subjects[0].Name) + assert.Equal(t, sa.Namespace, binding.Subjects[0].Namespace) } diff --git a/internal/app/app.go b/internal/app/app.go index 2ea365c8..c321f47a 100644 --- a/internal/app/app.go +++ b/internal/app/app.go @@ -908,10 +908,8 @@ func (a *App) Start(ctx context.Context) error { go a.startNodeRefresh(ctx) } - // Local-runtime topology reconcile into the per-agent store (under LocalAgent), - // so store-backed K8s/Swarm views reflect the local cluster too. Each is a - // no-op unless the matching runtime is active. - go a.startKubernetesReconcile(ctx) + // Local Swarm topology reconcile into the per-agent store (under LocalAgent); + // the Kubernetes one is wired with each runtime connection cycle. go a.startSwarmTopologyReconcile(ctx) // Sustained container downtime (opt-in via MAINTENANT_CONTAINER_DOWN_AFTER). diff --git a/internal/app/config.go b/internal/app/config.go index 3357459f..22a8557e 100644 --- a/internal/app/config.go +++ b/internal/app/config.go @@ -379,6 +379,7 @@ func ConfigFromEnv() Config { EnrollmentToken: os.Getenv("MAINTENANT_ENROLLMENT_TOKEN"), RuntimeOverride: os.Getenv("MAINTENANT_RUNTIME"), Label: os.Getenv("MAINTENANT_LABEL"), + NodeName: os.Getenv("MAINTENANT_NODE_NAME"), InsecureSkipVerify: parseTruthy(os.Getenv("MAINTENANT_GRPC_INSECURE_SKIP_TLS_VERIFY")), EmbeddedAgent: parseTruthy(os.Getenv("MAINTENANT_EMBEDDED_AGENT")), } diff --git a/internal/app/config_test.go b/internal/app/config_test.go index 9965c4d5..5df75494 100644 --- a/internal/app/config_test.go +++ b/internal/app/config_test.go @@ -172,6 +172,12 @@ func TestValidateProxiesRefusesGarbage(t *testing.T) { } } +func TestConfigFromEnv_NodeName(t *testing.T) { + t.Setenv("MAINTENANT_NODE_NAME", "worker-2") + + assert.Equal(t, "worker-2", ConfigFromEnv().MultiHost.NodeName) +} + func TestConfigFromEnv_AgentSpoolDefaults(t *testing.T) { cfg := ConfigFromEnv() diff --git a/internal/app/lifecycle.go b/internal/app/lifecycle.go index 680fff78..fed7f4ff 100644 --- a/internal/app/lifecycle.go +++ b/internal/app/lifecycle.go @@ -293,12 +293,8 @@ func (a *App) startContainerDownCheck(ctx context.Context) { // startKubernetesReconcile periodically snapshots the server's own Kubernetes // runtime into the per-agent store under the LocalAgent id, so the store-backed // Workloads/Pods/Nodes views reflect the local cluster the same way they reflect -// remote agents. No-op unless the local runtime is Kubernetes. -func (a *App) startKubernetesReconcile(ctx context.Context) { - src, ok := a.rt.(kubernetes.SnapshotSource) - if !ok || a.k8sIngest == nil { - return - } +// remote agents. It runs for one connection cycle and stops when until closes. +func (a *App) startKubernetesReconcile(ctx context.Context, src kubernetes.SnapshotSource, until <-chan struct{}) { reconcile := func() { snap, err := kubernetes.SnapshotFromRuntime(ctx, src) if err != nil { @@ -316,6 +312,8 @@ func (a *App) startKubernetesReconcile(ctx context.Context) { select { case <-ctx.Done(): return + case <-until: + return case <-ticker.C: reconcile() } @@ -508,7 +506,11 @@ func (a *App) startSwarmRecheck(ctx context.Context) { // Appelée exactement une fois par cycle de connexion (la garde est le superviseur). func (a *App) wireContainerMonitoring(ctx context.Context) <-chan struct{} { a.reconcile(ctx) - return a.startEventStream(ctx) + streamDone := a.startEventStream(ctx) + if src, ok := a.rt.(kubernetes.SnapshotSource); ok { + go a.startKubernetesReconcile(ctx, src, streamDone) + } + return streamDone } // broadcastRuntimeAvailability diffuse l'état runtime courant via SSE. diff --git a/internal/app/lifecycle_test.go b/internal/app/lifecycle_test.go index 6820b489..5f2c2cc8 100644 --- a/internal/app/lifecycle_test.go +++ b/internal/app/lifecycle_test.go @@ -4,6 +4,7 @@ import ( "context" "io" "log/slog" + "os" "path/filepath" "testing" "time" @@ -50,6 +51,40 @@ func TestSupervisor_DegradedThenHTTPUp(t *testing.T) { } } +// A kubeconfig whose cluster does not answer starts the server degraded, and it +// must stay up until shutdown instead of reading from a client never connected. +func TestStart_KubernetesUnreachable(t *testing.T) { + tmpDir := t.TempDir() + kubeconfig := filepath.Join(tmpDir, "kubeconfig") + cfg := "apiVersion: v1\nkind: Config\n" + + "clusters:\n- name: gone\n cluster:\n server: https://127.0.0.1:1\n" + + "users:\n- name: gone\n user:\n token: gone\n" + + "contexts:\n- name: gone\n context:\n cluster: gone\n user: gone\n" + + "current-context: gone\n" + require.NoError(t, os.WriteFile(kubeconfig, []byte(cfg), 0o600)) + t.Setenv("MAINTENANT_RUNTIME", "kubernetes") + t.Setenv("KUBERNETES_SERVICE_HOST", "") + t.Setenv("KUBECONFIG", kubeconfig) + + a, err := app.New(app.Config{ + DBPath: filepath.Join(tmpDir, "test.db"), + Addr: "127.0.0.1:0", + }, slog.New(slog.NewTextHandler(io.Discard, nil))) + require.NoError(t, err) + + ctx, cancel := context.WithTimeout(context.Background(), 3*time.Second) + defer cancel() + done := make(chan error, 1) + go func() { done <- a.Start(ctx) }() + + select { + case err := <-done: + assert.NoError(t, err) + case <-time.After(10 * time.Second): + t.Fatal("Start() did not return after context cancellation") + } +} + // TestSupervisor_NoGoroutineLeak (T023.c): flapping detection via goroutine count. // This is a smoke test: we start and stop the app multiple times to ensure // no goroutine growth. Exact goroutine count is not asserted (too fragile); diff --git a/internal/kubernetes/rbac.go b/internal/kubernetes/rbac.go new file mode 100644 index 00000000..0dece02c --- /dev/null +++ b/internal/kubernetes/rbac.go @@ -0,0 +1,18 @@ +// Copyright 2026 Benjamin Touchard (kOlapsis) +// SPDX-License-Identifier: Apache-2.0 + +package kubernetes + +import rbacv1 "k8s.io/api/rbac/v1" + +// ReadRules returns the read-only ClusterRole rules covering every API call the runtime makes. +func ReadRules() []rbacv1.PolicyRule { + read := []string{"get", "list", "watch"} + return []rbacv1.PolicyRule{ + {APIGroups: []string{""}, Resources: []string{"namespaces", "nodes", "pods", "events", "services"}, Verbs: read}, + {APIGroups: []string{""}, Resources: []string{"pods/log"}, Verbs: []string{"get"}}, + {APIGroups: []string{"apps"}, Resources: []string{"deployments", "statefulsets", "daemonsets"}, Verbs: read}, + {APIGroups: []string{"batch"}, Resources: []string{"jobs"}, Verbs: read}, + {APIGroups: []string{"metrics.k8s.io"}, Resources: []string{"pods", "nodes"}, Verbs: []string{"get", "list"}}, + } +} diff --git a/internal/kubernetes/rbac_test.go b/internal/kubernetes/rbac_test.go new file mode 100644 index 00000000..13bcd161 --- /dev/null +++ b/internal/kubernetes/rbac_test.go @@ -0,0 +1,219 @@ +// Copyright 2026 Benjamin Touchard (kOlapsis) +// SPDX-License-Identifier: Apache-2.0 + +package kubernetes + +import ( + "bytes" + "errors" + "io" + "os" + "path/filepath" + "regexp" + "sort" + "strings" + "testing" + + rbacv1 "k8s.io/api/rbac/v1" + metav1 "k8s.io/apimachinery/pkg/apis/meta/v1" + "k8s.io/apimachinery/pkg/util/yaml" +) + +var ( + clientCallRe = regexp.MustCompile(`\.(CoreV1|AppsV1|BatchV1|MetricsV1beta1)\(\)\.(\w+)\([^()]*\)\.(\w+)\(`) + informerCallRe = regexp.MustCompile(`\.(Core|Apps|Batch)\(\)\.V1\(\)\.(\w+)\(\)\.Informer\(\)`) +) + +var apiGroupOf = map[string]string{ + "CoreV1": "", "AppsV1": "apps", "BatchV1": "batch", "MetricsV1beta1": "metrics.k8s.io", + "Core": "", "Apps": "apps", "Batch": "batch", +} + +var metricsResource = map[string]string{"podmetricses": "pods", "nodemetricses": "nodes"} + +func grant(group, resource, verb string) string { + return group + "|" + resource + "|" + verb +} + +func grantsOf(rules []rbacv1.PolicyRule) map[string]bool { + out := map[string]bool{} + for _, r := range rules { + for _, g := range r.APIGroups { + for _, res := range r.Resources { + for _, v := range r.Verbs { + out[grant(g, res, v)] = true + } + } + } + } + return out +} + +func sortedKeys(m map[string]bool) []string { + keys := make([]string, 0, len(m)) + for k := range m { + keys = append(keys, k) + } + sort.Strings(keys) + return keys +} + +// apiCallsInPackage reads the package sources and returns every (group, resource, verb) they call. +func apiCallsInPackage(t *testing.T) map[string]bool { + t.Helper() + files, err := filepath.Glob("*.go") + if err != nil { + t.Fatal(err) + } + calls := map[string]bool{} + for _, f := range files { + if strings.HasSuffix(f, "_test.go") { + continue + } + src, err := os.ReadFile(f) + if err != nil { + t.Fatal(err) + } + for _, m := range clientCallRe.FindAllStringSubmatch(string(src), -1) { + group, resource := apiGroupOf[m[1]], strings.ToLower(m[2]) + if r, ok := metricsResource[resource]; ok { + resource = r + } + switch m[3] { + case "Get", "List", "Watch": + calls[grant(group, resource, strings.ToLower(m[3]))] = true + case "GetLogs": + calls[grant(group, resource+"/log", "get")] = true + default: + t.Errorf("%s: %s().%s().%s is not a read call", f, m[1], m[2], m[3]) + } + } + for _, m := range informerCallRe.FindAllStringSubmatch(string(src), -1) { + group, resource := apiGroupOf[m[1]], strings.ToLower(m[2]) + calls[grant(group, resource, "list")] = true + calls[grant(group, resource, "watch")] = true + } + } + return calls +} + +func TestReadRulesCoverEveryAPICall(t *testing.T) { + calls := apiCallsInPackage(t) + if !calls[grant("", "pods", "list")] || !calls[grant("metrics.k8s.io", "nodes", "get")] { + t.Fatalf("source scan found too little, the patterns no longer match the client calls: %v", sortedKeys(calls)) + } + granted := grantsOf(ReadRules()) + for _, c := range sortedKeys(calls) { + if !granted[c] { + t.Errorf("the runtime calls %s but ReadRules does not grant it", c) + } + } +} + +func TestReadRulesAreReadOnly(t *testing.T) { + for _, r := range ReadRules() { + for _, v := range r.Verbs { + if v != "get" && v != "list" && v != "watch" { + t.Errorf("rule %v grants %q", r.Resources, v) + } + } + } +} + +type manifestObject struct { + Kind string `json:"kind"` + Metadata metav1.ObjectMeta `json:"metadata"` + Rules []rbacv1.PolicyRule `json:"rules"` + Subjects []rbacv1.Subject `json:"subjects"` +} + +func decodeManifests(t *testing.T, data []byte) []manifestObject { + t.Helper() + dec := yaml.NewYAMLOrJSONDecoder(bytes.NewReader(data), 4096) + var objs []manifestObject + for { + var obj manifestObject + err := dec.Decode(&obj) + if errors.Is(err, io.EOF) { + return objs + } + if err != nil { + t.Fatalf("decode manifest: %v", err) + } + if obj.Kind != "" { + objs = append(objs, obj) + } + } +} + +func readManifest(t *testing.T, path string) []byte { + t.Helper() + data, err := os.ReadFile(filepath.Join("..", "..", path)) + if err != nil { + t.Fatal(err) + } + return data +} + +// withoutTemplateLines drops the Helm action lines so the chart's static rules parse as YAML. +func withoutTemplateLines(data []byte) []byte { + var out [][]byte + for _, line := range bytes.Split(data, []byte("\n")) { + if !bytes.Contains(line, []byte("{{")) { + out = append(out, line) + } + } + return bytes.Join(out, []byte("\n")) +} + +func TestServerManifestsGrantExactlyReadRules(t *testing.T) { + sources := map[string][]byte{ + "deploy/kubernetes/rbac.yaml": readManifest(t, "deploy/kubernetes/rbac.yaml"), + "deploy/helm/maintenant/templates/rbac.yaml": withoutTemplateLines(readManifest(t, "deploy/helm/maintenant/templates/rbac.yaml")), + } + want := grantsOf(ReadRules()) + for path, data := range sources { + var roles int + for _, obj := range decodeManifests(t, data) { + if obj.Kind != "ClusterRole" { + continue + } + roles++ + got := grantsOf(obj.Rules) + for _, g := range sortedKeys(want) { + if !got[g] { + t.Errorf("%s: ClusterRole misses %s", path, g) + } + } + for _, g := range sortedKeys(got) { + if !want[g] { + t.Errorf("%s: ClusterRole grants %s, which the runtime never uses", path, g) + } + } + } + if roles != 1 { + t.Errorf("%s: %d ClusterRole, want 1", path, roles) + } + } +} + +func TestPlainManifestsLiveInTheMaintenantNamespace(t *testing.T) { + const namespace = "maintenant" + clusterScoped := map[string]bool{"ClusterRole": true, "ClusterRoleBinding": true, "Namespace": true} + for _, path := range []string{"deploy/kubernetes/deployment.yaml", "deploy/kubernetes/rbac.yaml"} { + for _, obj := range decodeManifests(t, readManifest(t, path)) { + if clusterScoped[obj.Kind] { + if obj.Metadata.Namespace != "" { + t.Errorf("%s: cluster-scoped %s %q carries a namespace", path, obj.Kind, obj.Metadata.Name) + } + } else if obj.Metadata.Namespace != namespace { + t.Errorf("%s: %s %q is in namespace %q, want %q", path, obj.Kind, obj.Metadata.Name, obj.Metadata.Namespace, namespace) + } + for _, s := range obj.Subjects { + if s.Kind == "ServiceAccount" && s.Namespace != namespace { + t.Errorf("%s: %s binds ServiceAccount %s/%s, want namespace %q", path, obj.Metadata.Name, s.Namespace, s.Name, namespace) + } + } + } + } +} diff --git a/internal/kubernetes/runtime.go b/internal/kubernetes/runtime.go index 33878116..dfc0c389 100644 --- a/internal/kubernetes/runtime.go +++ b/internal/kubernetes/runtime.go @@ -13,6 +13,7 @@ import ( "time" cmodel "github.com/kolapsis/maintenant/internal/container" + "github.com/kolapsis/maintenant/internal/retry" "github.com/kolapsis/maintenant/internal/runtime" metav1 "k8s.io/apimachinery/pkg/apis/meta/v1" "k8s.io/client-go/informers" @@ -66,12 +67,35 @@ func NewRuntime(logger *slog.Logger, nsFilter *NamespaceFilter) (*Runtime, error }, nil } +const ( + connectInitialBackoff = 1 * time.Second + connectMaxBackoff = 30 * time.Second +) + +// Connect retries until the API server answers or ctx is cancelled, but returns a configuration error at once. func (r *Runtime) Connect(ctx context.Context) error { config, err := buildConfig() if err != nil { return fmt.Errorf("kubernetes config: %w", err) } + b := retry.New(connectInitialBackoff, connectMaxBackoff, 0) + for { + err := r.connect(ctx, config) + if err == nil { + return nil + } + delay := b.Next() + r.logger.Warn("kubernetes connection failed, retrying", "error", err, "retry_in", delay) + select { + case <-ctx.Done(): + return ctx.Err() + case <-time.After(delay): + } + } +} + +func (r *Runtime) connect(ctx context.Context, config *rest.Config) error { clientset, err := k8s.NewForConfig(config) if err != nil { return fmt.Errorf("kubernetes clientset: %w", err) @@ -82,9 +106,7 @@ func (r *Runtime) Connect(ctx context.Context) error { r.logger.Warn("metrics-server client failed; resource metrics will be unavailable", "error", err) } - // Verify connectivity. - _, err = clientset.Discovery().ServerVersion() - if err != nil { + if _, err := clientset.Discovery().RESTClient().Get().AbsPath("/version").DoRaw(ctx); err != nil { return fmt.Errorf("kubernetes connectivity check failed: %w", err) } @@ -117,10 +139,15 @@ func (r *Runtime) Connect(ctx context.Context) error { return nil } +// TryConnect makes a single attempt bounded to a few seconds. func (r *Runtime) TryConnect(ctx context.Context) error { + config, err := buildConfig() + if err != nil { + return fmt.Errorf("kubernetes config: %w", err) + } tctx, cancel := context.WithTimeout(ctx, 3*time.Second) defer cancel() - return r.Connect(tctx) + return r.connect(tctx, config) } func buildConfig() (*rest.Config, error) { diff --git a/internal/kubernetes/runtime_test.go b/internal/kubernetes/runtime_test.go new file mode 100644 index 00000000..8ccf4621 --- /dev/null +++ b/internal/kubernetes/runtime_test.go @@ -0,0 +1,90 @@ +// Copyright 2026 Benjamin Touchard (kOlapsis) +// SPDX-License-Identifier: Apache-2.0 + +package kubernetes + +import ( + "context" + "io" + "log/slog" + "net/http" + "net/http/httptest" + "os" + "path/filepath" + "sync/atomic" + "testing" + "time" + + "github.com/stretchr/testify/assert" + "github.com/stretchr/testify/require" +) + +// fakeAPIServer answers /version with 503 for the first failures calls, then with a version. +func fakeAPIServer(t *testing.T, failures int32) *atomic.Int32 { + t.Helper() + var hits atomic.Int32 + srv := httptest.NewServer(http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) { + if r.URL.Path != "/version" { + http.NotFound(w, r) + return + } + if hits.Add(1) <= failures { + w.WriteHeader(http.StatusServiceUnavailable) + return + } + _, _ = w.Write([]byte(`{"major":"1","minor":"30","gitVersion":"v1.30.0"}`)) + })) + t.Cleanup(srv.Close) + + kubeconfig := filepath.Join(t.TempDir(), "kubeconfig") + cfg := "apiVersion: v1\nkind: Config\n" + + "clusters:\n- name: fake\n cluster:\n server: " + srv.URL + "\n" + + "users:\n- name: fake\n user:\n token: fake\n" + + "contexts:\n- name: fake\n context:\n cluster: fake\n user: fake\n" + + "current-context: fake\n" + require.NoError(t, os.WriteFile(kubeconfig, []byte(cfg), 0o600)) + t.Setenv("KUBERNETES_SERVICE_HOST", "") + t.Setenv("KUBECONFIG", kubeconfig) + return &hits +} + +func newTestRuntime(t *testing.T) *Runtime { + t.Helper() + r, err := NewRuntime(slog.New(slog.NewTextHandler(io.Discard, nil)), NewNamespaceFilter("", "")) + require.NoError(t, err) + t.Cleanup(func() { _ = r.Close() }) + return r +} + +func TestConnect_RetriesUntilTheAPIServerAnswers(t *testing.T) { + hits := fakeAPIServer(t, 1) + r := newTestRuntime(t) + + ctx, cancel := context.WithTimeout(context.Background(), 10*time.Second) + defer cancel() + require.NoError(t, r.Connect(ctx)) + + assert.True(t, r.IsConnected()) + assert.Equal(t, int32(2), hits.Load()) +} + +func TestConnect_StopsWithItsContext(t *testing.T) { + fakeAPIServer(t, 1<<30) + r := newTestRuntime(t) + + ctx, cancel := context.WithTimeout(context.Background(), 1500*time.Millisecond) + defer cancel() + + assert.ErrorIs(t, r.Connect(ctx), context.DeadlineExceeded) + assert.False(t, r.IsConnected()) +} + +func TestTryConnect_MakesASingleAttempt(t *testing.T) { + hits := fakeAPIServer(t, 1) + r := newTestRuntime(t) + + require.Error(t, r.TryConnect(context.Background())) + + assert.False(t, r.IsConnected()) + assert.Equal(t, int32(1), hits.Load()) +} diff --git a/scripts/ci-changed-paths.sh b/scripts/ci-changed-paths.sh index 88c41c04..0653d4c2 100755 --- a/scripts/ci-changed-paths.sh +++ b/scripts/ci-changed-paths.sh @@ -43,12 +43,16 @@ if [ -z "$reason" ]; then list=$(mktemp) printf '%s\n' "$files" >"$list" - # A path can only fall in one bucket, so order matters: the inert list - # wins over everything, then the frontend, then the backend. What no - # pattern claims is unknown, and unknown means run it all. + # A path can only fall in one bucket, so order matters: the Kubernetes + # manifests come first, then the inert list, then the frontend, then the + # backend. What no pattern claims is unknown, and unknown means run it all. while read -r f; do [ -n "$f" ] || continue case "$f" in + # The Go tests of internal/kubernetes check their RBAC and namespaces. + deploy/kubernetes/* | deploy/helm/*) + backend=true + ;; # Nothing in CI reads these. The docs have their own workflow. docs/* | mkdocs.yml | deploy/* | *.md | LICENSE | .env.example | .gitignore) ;; frontend/*) From 3c285645fc99b0178b0504af131d97a3738f5bb2 Mon Sep 17 00:00:00 2001 From: Benjamin Date: Wed, 30 Sep 2026 19:03:51 +0200 Subject: [PATCH 05/54] fix(mcp): serve /mcp behind a reverse proxy on the same host On a loopback listener the MCP SDK refuses any request whose Host header is not loopback, so a local nginx forwarding the public Host got 403 "invalid Host header" on every /mcp call. When MCP OAuth is configured, the bearer token already authenticates each request, so the SDK's localhost protection is turned off in that case only. The unauthenticated mode keeps it. --- internal/app/http.go | 8 +++-- internal/app/http_wiring_test.go | 62 ++++++++++++++++++++++++++++++++ 2 files changed, 68 insertions(+), 2 deletions(-) diff --git a/internal/app/http.go b/internal/app/http.go index ce67ec82..175bf9df 100644 --- a/internal/app/http.go +++ b/internal/app/http.go @@ -213,11 +213,15 @@ func (a *App) buildHTTPServer() *http.Server { topMux := http.NewServeMux() if a.cfg.MCP.Enabled && !a.cfg.DemoMode { + oauthConfigured := a.cfg.MCP.ClientID != "" && a.cfg.MCP.ClientSecret != "" mcpHandler := unbufferedStream(gomcp.NewStreamableHTTPHandler(func(_ *http.Request) *gomcp.Server { return a.mcpServer - }, nil)) + }, &gomcp.StreamableHTTPOptions{ + // A reverse proxy on the same host forwards the public Host, which the SDK refuses on a loopback listener. + DisableLocalhostProtection: oauthConfigured, + })) - if a.cfg.MCP.ClientID != "" && a.cfg.MCP.ClientSecret != "" { + if oauthConfigured { a.warnWeakMCPClientSecret() mcpOAuthStore := store.NewMCPOAuthStore(a.db) oauthSrv := mcpoauth.NewOAuthServer(mcpoauth.Config{ diff --git a/internal/app/http_wiring_test.go b/internal/app/http_wiring_test.go index eae13969..5a0ded0c 100644 --- a/internal/app/http_wiring_test.go +++ b/internal/app/http_wiring_test.go @@ -12,11 +12,15 @@ import ( "net/url" "strings" "testing" + "time" "github.com/stretchr/testify/assert" "github.com/stretchr/testify/require" v1 "github.com/kolapsis/maintenant/internal/api/v1" + mcpoauth "github.com/kolapsis/maintenant/internal/mcp/oauth" + "github.com/kolapsis/maintenant/internal/store" + "github.com/kolapsis/maintenant/internal/uid" ) const wiringWebhookBody = `{"name":"hook","url":"https://8.8.8.8/hook","event_types":["*"]}` @@ -122,6 +126,64 @@ func TestHTTPServer_MCPStreamIsNotBuffered(t *testing.T) { assert.Equal(t, "no", rec.Header().Get("X-Accel-Buffering")) } +func postInitializeViaLocalProxy(t *testing.T, a *App, bearer string) (*http.Response, string) { + t.Helper() + srv := httptest.NewServer(a.srv.Handler) + t.Cleanup(srv.Close) + + req, err := http.NewRequest(http.MethodPost, srv.URL+"/mcp", strings.NewReader(initializeRequest)) + require.NoError(t, err) + req.Host = "maintenant.example.com" + req.Header.Set("Content-Type", "application/json") + req.Header.Set("Accept", "application/json, text/event-stream") + if bearer != "" { + req.Header.Set("Authorization", "Bearer "+bearer) + } + resp, err := srv.Client().Do(req) + require.NoError(t, err) + t.Cleanup(func() { _ = resp.Body.Close() }) + body, err := io.ReadAll(resp.Body) + require.NoError(t, err) + return resp, string(body) +} + +func TestHTTPServer_MCPWithOAuthServesALocalReverseProxy(t *testing.T) { + a, _ := newTestApp(t, func(c *Config) { + c.MCP.Enabled = true + c.MCP.ClientID = "claude" + c.MCP.ClientSecret = strings.Repeat("0f", 32) + c.BaseURL = "https://maintenant.example.com" + }) + token, hash := mcpoauth.GenerateToken() + require.NoError(t, store.NewMCPOAuthStore(a.db).StoreToken(context.Background(), &mcpoauth.MCPOAuthToken{ + TokenHash: hash, + TokenType: "access", + ClientID: "claude", + ExpiresAt: time.Now().Add(time.Hour), + FamilyID: uid.New(), + CreatedAt: time.Now(), + })) + + resp, body := postInitializeViaLocalProxy(t, a, token) + assert.Equal(t, http.StatusOK, resp.StatusCode, body) + assert.Equal(t, "text/event-stream", resp.Header.Get("Content-Type")) + + resp, _ = postInitializeViaLocalProxy(t, a, "") + assert.Equal(t, http.StatusUnauthorized, resp.StatusCode, "the bearer token still guards /mcp") +} + +func TestHTTPServer_MCPWithoutAuthKeepsTheLocalhostProtection(t *testing.T) { + a, _ := newTestApp(t, func(c *Config) { + c.MCP.Enabled = true + c.MCP.AllowUnauthenticated = true + }) + + resp, body := postInitializeViaLocalProxy(t, a, "") + + assert.Equal(t, http.StatusForbidden, resp.StatusCode) + assert.Contains(t, body, "invalid Host header") +} + func TestHTTPServer_WarnsAboutAShortMCPClientSecret(t *testing.T) { for name, tc := range map[string]struct { secret string From b41ab8fe4d46c6a73efd4c5bdf4f63c04f90dc0b Mon Sep 17 00:00:00 2001 From: Benjamin Date: Wed, 30 Sep 2026 19:03:34 +0200 Subject: [PATCH 06/54] fix(status): send status page emails through the environment SMTP server The subscriber service and the notifier were built with a nil mailer, so no confirmation or notification email ever left. Both now use the SMTP client of the email channel, built from MAINTENANT_SMTP_*. Mails are plain text, and the SMTP exchange honours its context with a 30 s ceiling. Manual incidents (API and MCP) now notify confirmed subscribers on creation, on each update and on resolution, through the same announce path as automatic incidents. A resolving update sends one resolution mail, not an update too. The in-memory status page SMTP settings are removed (GET/PUT /api/v1/status/smtp, form, types). POST /api/v1/status/smtp/test takes a recipient and uses the environment configuration. /status/api exposes subscriptions_enabled and the public page shows the subscription form only then. POST /status/subscribe answers JSON errors (subscriptions_unavailable, confirmation_failed, ...) and its 429 matches the other limiters, with Retry-After. --- .../components/status/StatusSmtpConfig.vue | 195 +++----------- .../components/status/SubscribersPanel.vue | 8 + .../status/__tests__/StatusSmtpConfig.spec.ts | 43 +++ frontend/src/locales/status-page/en.ts | 6 + frontend/src/locales/status-page/fr.ts | 6 + frontend/src/pages/PublicStatusPage.vue | 79 ++++++ .../pages/__tests__/PublicStatusPage.spec.ts | 83 ++++++ frontend/src/services/statusApi.ts | 30 +-- go.mod | 1 - go.sum | 2 - internal/api/v1/quota_tiers_test.go | 2 +- internal/api/v1/router.go | 8 +- internal/api/v1/status_admin.go | 114 ++------ internal/api/v1/status_mail_test.go | 238 +++++++++++++++++ internal/api/v1/status_tiers_test.go | 5 +- internal/app/app.go | 26 +- internal/app/status_mail_wiring_test.go | 37 +++ internal/commercial/channels/smtp.go | 18 +- internal/commercial/channels/smtp_test.go | 75 ++++++ internal/commercial/statuspage/factory.go | 16 +- internal/commercial/statuspage/incidents.go | 19 +- internal/commercial/statuspage/notifier.go | 25 +- .../commercial/statuspage/notifier_test.go | 158 +++++++++++ internal/commercial/statuspage/smtp.go | 107 -------- internal/extpoint/extpoint.go | 5 +- internal/mcp/server.go | 7 + internal/mcp/tools_write.go | 6 + internal/mcp/tools_write_test.go | 41 +++ internal/ratelimit/limiter_test.go | 31 +++ internal/ratelimit/middleware.go | 27 +- internal/status/announce.go | 58 ++++ internal/status/announce_test.go | 249 ++++++++++++++++++ internal/status/extension.go | 4 +- internal/status/handler.go | 51 +++- internal/status/handler_test.go | 204 +++++++++++++- internal/status/model.go | 18 -- internal/status/service.go | 23 +- internal/status/subscriber.go | 45 +++- 38 files changed, 1547 insertions(+), 523 deletions(-) create mode 100644 frontend/src/commercial/components/status/__tests__/StatusSmtpConfig.spec.ts create mode 100644 frontend/src/pages/__tests__/PublicStatusPage.spec.ts create mode 100644 internal/api/v1/status_mail_test.go create mode 100644 internal/app/status_mail_wiring_test.go create mode 100644 internal/commercial/channels/smtp_test.go create mode 100644 internal/commercial/statuspage/notifier_test.go delete mode 100644 internal/commercial/statuspage/smtp.go create mode 100644 internal/status/announce.go create mode 100644 internal/status/announce_test.go diff --git a/frontend/src/commercial/components/status/StatusSmtpConfig.vue b/frontend/src/commercial/components/status/StatusSmtpConfig.vue index 6d3e1fc1..fc660e77 100644 --- a/frontend/src/commercial/components/status/StatusSmtpConfig.vue +++ b/frontend/src/commercial/components/status/StatusSmtpConfig.vue @@ -5,186 +5,65 @@ --> diff --git a/frontend/src/commercial/components/status/SubscribersPanel.vue b/frontend/src/commercial/components/status/SubscribersPanel.vue index 5a49ccd7..d612f2a5 100644 --- a/frontend/src/commercial/components/status/SubscribersPanel.vue +++ b/frontend/src/commercial/components/status/SubscribersPanel.vue @@ -6,8 +6,11 @@ - - - - diff --git a/frontend/src/components/TriggerList.vue b/frontend/src/components/TriggerList.vue index 714718cc..11294bfd 100644 --- a/frontend/src/components/TriggerList.vue +++ b/frontend/src/components/TriggerList.vue @@ -36,7 +36,6 @@ function summarizeFilter(t: AlertTrigger): string { if (t.filter_severities) parts.push(`severity: ${t.filter_severities}`) if (t.filter_sources) parts.push(`source: ${t.filter_sources}`) if (t.filter_scopes) parts.push(`scope: ${t.filter_scopes}`) - if (t.filter_tags) parts.push(`tags: ${t.filter_tags}`) if (parts.length === 0) return 'matches all alerts' return parts.join(' · ') } diff --git a/frontend/src/components/TriggerManager.vue b/frontend/src/components/TriggerManager.vue index 68c036aa..41165e30 100644 --- a/frontend/src/components/TriggerManager.vue +++ b/frontend/src/components/TriggerManager.vue @@ -58,7 +58,6 @@ async function handleToggleEnabled(t: AlertTrigger) { filter_severities: t.filter_severities, filter_sources: t.filter_sources, filter_scopes: t.filter_scopes, - filter_tags: t.filter_tags, enabled: !t.enabled, notify_on_resolve: t.notify_on_resolve, channel_ids: t.channel_ids, diff --git a/frontend/src/components/__tests__/TriggerEditor.spec.ts b/frontend/src/components/__tests__/TriggerEditor.spec.ts new file mode 100644 index 00000000..86ff28f6 --- /dev/null +++ b/frontend/src/components/__tests__/TriggerEditor.spec.ts @@ -0,0 +1,72 @@ +// Copyright 2026 Benjamin Touchard (kOlapsis) +// SPDX-License-Identifier: Apache-2.0 + +import { describe, it, expect, vi, beforeEach } from 'vitest' +import { mount, flushPromises, RouterLinkStub } from '@vue/test-utils' +import type { AlertTrigger } from '@/types/triggers' +import TriggerEditor from '@/components/TriggerEditor.vue' + +const update = vi.fn() + +vi.mock('@/stores/triggers', () => ({ + useTriggersStore: () => ({ create: vi.fn(), update }), +})) + +vi.mock('@/stores/channels', () => ({ + useChannelsStore: () => ({ channels: [{ id: 'c1', name: 'ops', type: 'webhook', enabled: true }] }), +})) + +vi.mock('@/composables/useEdition', () => ({ + useEdition: () => ({ hasFeature: () => true, requiredEditionFor: () => 'pro' }), +})) + +const trigger: AlertTrigger = { + id: 't1', + name: 'Critical containers', + filter_severities: 'critical', + filter_sources: 'container', + filter_scopes: 'container:42', + enabled: true, + notify_on_resolve: true, + channel_ids: ['c1'], + created_at: '', + updated_at: '', +} + +describe('TriggerEditor', () => { + beforeEach(() => { + update.mockReset() + }) + + it('offers no tag filter', () => { + const wrapper = mount(TriggerEditor, { + props: { trigger }, + global: { stubs: { RouterLink: RouterLinkStub } }, + }) + expect(wrapper.text()).not.toMatch(/tag/i) + }) + + it('saves the filters it shows and nothing else', async () => { + const wrapper = mount(TriggerEditor, { + props: { trigger }, + global: { stubs: { RouterLink: RouterLinkStub } }, + }) + const save = wrapper.findAll('button').find((b) => b.text().includes('Save changes')) + expect(save).toBeDefined() + await save!.trigger('click') + await flushPromises() + + expect(update).toHaveBeenCalledOnce() + const [id, req] = update.mock.calls[0]! + expect(id).toBe('t1') + expect(req).toEqual({ + name: 'Critical containers', + filter_severities: 'critical', + filter_sources: 'container', + filter_scopes: 'container:42', + enabled: true, + notify_on_resolve: true, + channel_ids: ['c1'], + }) + }) +}) diff --git a/frontend/src/services/containerApi.ts b/frontend/src/services/containerApi.ts index 769008a0..58169bf8 100644 --- a/frontend/src/services/containerApi.ts +++ b/frontend/src/services/containerApi.ts @@ -23,7 +23,6 @@ export interface Container { is_ignored: boolean alert_severity: string restart_threshold: number - alert_channels?: string archived: boolean first_seen_at: string last_state_change_at: string diff --git a/frontend/src/types/triggers.ts b/frontend/src/types/triggers.ts index f2ddf79f..e45a6d8e 100644 --- a/frontend/src/types/triggers.ts +++ b/frontend/src/types/triggers.ts @@ -7,7 +7,6 @@ export interface AlertTrigger { filter_severities: string filter_sources: string filter_scopes: string - filter_tags: string enabled: boolean notify_on_resolve: boolean channel_ids: string[] @@ -20,7 +19,6 @@ export interface TriggerRequest { filter_severities: string filter_sources: string filter_scopes: string - filter_tags: string enabled: boolean notify_on_resolve: boolean channel_ids: string[] diff --git a/internal/alert/alerts.go b/internal/alert/alerts.go index 27bcda40..d9129997 100644 --- a/internal/alert/alerts.go +++ b/internal/alert/alerts.go @@ -41,7 +41,6 @@ type RestartAlert struct { RestartCount int Threshold int Severity container.AlertSeverity - Channels string Timestamp time.Time AgentID string } @@ -72,7 +71,6 @@ func (d *RestartDetector) Check(ctx context.Context, c *container.Container) (in RestartCount: count, Threshold: c.RestartThreshold, Severity: c.AlertSeverity, - Channels: c.AlertChannels, Timestamp: time.Now(), AgentID: c.AgentID, }, nil @@ -85,7 +83,6 @@ type HealthAlert struct { PreviousHealth *container.HealthStatus NewHealth container.HealthStatus Severity container.AlertSeverity - Channels string Timestamp time.Time } @@ -102,7 +99,6 @@ func CheckHealthTransition(c *container.Container, previousHealth *container.Hea PreviousHealth: previousHealth, NewHealth: newHealth, Severity: c.AlertSeverity, - Channels: c.AlertChannels, Timestamp: time.Now(), } } diff --git a/internal/alert/engine.go b/internal/alert/engine.go index 6ebd4f5f..6aeaac71 100644 --- a/internal/alert/engine.go +++ b/internal/alert/engine.go @@ -639,9 +639,6 @@ func (e *Engine) enqueueDelivery(ctx context.Context, ch *NotificationChannel, a // matchesTrigger reports whether an alert satisfies all of a trigger's // non-empty filters (AND between fields, OR within a CSV field). An empty // filter matches everything. -// -// FilterTags is treated as no-op for now: Alert does not yet expose tags. -// The match is enforced via filter_severities, filter_sources and filter_scopes. func matchesTrigger(t *AlertTrigger, a *Alert) bool { if a.Status == StatusResolved && !t.NotifyOnResolve { return false diff --git a/internal/alert/escalation/model.go b/internal/alert/escalation/model.go index daa521aa..e172c44b 100644 --- a/internal/alert/escalation/model.go +++ b/internal/alert/escalation/model.go @@ -23,7 +23,6 @@ type Policy struct { type Filters struct { Severities []string `json:"severities"` Scopes []Scope `json:"scopes"` - Tags []string `json:"tags"` } // Scope identifies a specific monitored entity. diff --git a/internal/alert/model.go b/internal/alert/model.go index ae5c70fb..ae2d1c87 100644 --- a/internal/alert/model.go +++ b/internal/alert/model.go @@ -112,14 +112,13 @@ type NotificationChannel struct { // AlertTrigger is a routing rule that maps an alert filter to one or more channels. // Filters are stored as CSV strings; an empty filter matches anything. // Filters are combined in AND between fields, OR within a field. -// FilterScopes and FilterTags are Pro-only (gated at the handler level). +// FilterScopes is Pro-only (gated at the handler level). type AlertTrigger struct { ID string `json:"id"` Name string `json:"name"` FilterSeverities string `json:"filter_severities"` FilterSources string `json:"filter_sources"` FilterScopes string `json:"filter_scopes"` - FilterTags string `json:"filter_tags"` Enabled bool `json:"enabled"` NotifyOnResolve bool `json:"notify_on_resolve"` ChannelIDs []string `json:"channel_ids"` diff --git a/internal/api/v1/alert_triggers.go b/internal/api/v1/alert_triggers.go index 3e5b5c2a..c173b64b 100644 --- a/internal/api/v1/alert_triggers.go +++ b/internal/api/v1/alert_triggers.go @@ -37,7 +37,6 @@ type triggerInput struct { FilterSeverities string `json:"filter_severities"` FilterSources string `json:"filter_sources"` FilterScopes string `json:"filter_scopes"` - FilterTags string `json:"filter_tags"` Enabled *bool `json:"enabled"` NotifyOnResolve *bool `json:"notify_on_resolve"` ChannelIDs []string `json:"channel_ids"` @@ -88,8 +87,7 @@ func (h *AlertTriggerHandler) HandleCreateTrigger(w http.ResponseWriter, r *http WriteError(w, http.StatusBadRequest, "validation_failed", err.Error()) return } - if err := h.checkAdvancedFiltersGating(&input); err != nil { - WriteError(w, http.StatusForbidden, "edition_required", err.Error()) + if refuseAdvancedFilters(w, &input) { return } if err := h.checkChannelsExist(r, input.ChannelIDs); err != nil { @@ -111,7 +109,6 @@ func (h *AlertTriggerHandler) HandleCreateTrigger(w http.ResponseWriter, r *http FilterSeverities: input.FilterSeverities, FilterSources: input.FilterSources, FilterScopes: input.FilterScopes, - FilterTags: input.FilterTags, Enabled: enabled, NotifyOnResolve: notifyOnResolve, ChannelIDs: input.ChannelIDs, @@ -167,8 +164,7 @@ func (h *AlertTriggerHandler) HandleUpdateTrigger(w http.ResponseWriter, r *http WriteError(w, http.StatusBadRequest, "validation_failed", err.Error()) return } - if err := h.checkAdvancedFiltersGating(&input); err != nil { - WriteError(w, http.StatusForbidden, "edition_required", err.Error()) + if refuseAdvancedFilters(w, &input) { return } if err := h.checkChannelsExist(r, input.ChannelIDs); err != nil { @@ -189,7 +185,6 @@ func (h *AlertTriggerHandler) HandleUpdateTrigger(w http.ResponseWriter, r *http existing.FilterSeverities = input.FilterSeverities existing.FilterSources = input.FilterSources existing.FilterScopes = input.FilterScopes - existing.FilterTags = input.FilterTags existing.Enabled = enabled existing.NotifyOnResolve = notifyOnResolve existing.ChannelIDs = input.ChannelIDs @@ -256,16 +251,14 @@ func validateTriggerInput(t *triggerInput) error { return nil } -// checkAdvancedFiltersGating returns an error when CE is used with Pro-only filters. -func (h *AlertTriggerHandler) checkAdvancedFiltersGating(t *triggerInput) error { - if t.FilterScopes == "" && t.FilterTags == "" { - return nil - } - if extension.Allows(extension.CapAlertAdvancedFilters) { - return nil - } - return errors.New("advanced filters (scopes, tags) require the " + - titleEdition(extension.MinEdition(extension.CapAlertAdvancedFilters)) + " edition") +// refuseAdvancedFilters writes the edition refusal when a scope filter is set +// on an edition that does not open it, and reports whether it did. +func refuseAdvancedFilters(w http.ResponseWriter, t *triggerInput) bool { + if t.FilterScopes == "" || extension.Allows(extension.CapAlertAdvancedFilters) { + return false + } + refuseCapability(w, extension.CapAlertAdvancedFilters) + return true } // checkChannelsExist verifies that each channel_id points at an existing row. diff --git a/internal/api/v1/alert_triggers_test.go b/internal/api/v1/alert_triggers_test.go index 5ecfa6eb..fc22d18f 100644 --- a/internal/api/v1/alert_triggers_test.go +++ b/internal/api/v1/alert_triggers_test.go @@ -5,6 +5,7 @@ package v1 import ( "context" + "encoding/json" "fmt" "net/http" "net/http/httptest" @@ -247,23 +248,47 @@ func TestHandleCreateTrigger_ProFilterScopes_CommunityBlocked(t *testing.T) { req.Header.Set("Content-Type", "application/json") rec := httptest.NewRecorder() h.HandleCreateTrigger(rec, req) - assert.Equal(t, http.StatusForbidden, rec.Code) - assert.Contains(t, rec.Body.String(), "edition_required") + assertAdvancedFiltersRefusal(t, rec) } -func TestHandleCreateTrigger_ProFilterTags_CommunityBlocked(t *testing.T) { +func TestHandleUpdateTrigger_ProFilterScopes_CommunityBlocked(t *testing.T) { original := extension.CurrentEdition extension.CurrentEdition = func() extension.Edition { return extension.Community } defer func() { extension.CurrentEdition = original }() + h, ts := newTriggerHandler(true) + id := seedTrigger(t, ts, "Original") + body := `{"name":"Scoped","filter_scopes":"container:42","channel_ids":["1"]}` + req := httptest.NewRequest("PUT", "/api/v1/alert-triggers/1", strings.NewReader(body)) + req.Header.Set("Content-Type", "application/json") + req.SetPathValue("id", id) + rec := httptest.NewRecorder() + h.HandleUpdateTrigger(rec, req) + assertAdvancedFiltersRefusal(t, rec) + assert.Equal(t, "Original", ts.triggers[id].Name, "a refused update leaves the trigger alone") +} + +func assertAdvancedFiltersRefusal(t *testing.T, rec *httptest.ResponseRecorder) { + t.Helper() + assert.Equal(t, http.StatusForbidden, rec.Code) + var body ErrorResponse + require.NoError(t, json.Unmarshal(rec.Body.Bytes(), &body)) + assert.Equal(t, "EDITION_REQUIRED", body.Error.Code) + assert.Equal(t, string(extension.CapAlertAdvancedFilters), body.Error.Feature) + assert.Equal(t, string(extension.MinEdition(extension.CapAlertAdvancedFilters)), body.Error.RequiredEdition) +} + +func TestHandleCreateTrigger_HasNoTagFilter(t *testing.T) { h, _ := newTriggerHandler(true) - body := `{"name":"Tagged","filter_tags":"prod","channel_ids":["1"]}` + body := `{"name":"Plain","channel_ids":["1"]}` req := httptest.NewRequest("POST", "/api/v1/alert-triggers", strings.NewReader(body)) req.Header.Set("Content-Type", "application/json") rec := httptest.NewRecorder() h.HandleCreateTrigger(rec, req) - assert.Equal(t, http.StatusForbidden, rec.Code) - assert.Contains(t, rec.Body.String(), "edition_required") + require.Equal(t, http.StatusCreated, rec.Code) + var got map[string]any + require.NoError(t, json.Unmarshal(rec.Body.Bytes(), &got)) + assert.NotContains(t, got, "filter_tags") } func TestHandleCreateTrigger_ProFilterAllowed_Pro(t *testing.T) { diff --git a/internal/api/v1/alerts.go b/internal/api/v1/alerts.go index f20d08fc..af0fe7f8 100644 --- a/internal/api/v1/alerts.go +++ b/internal/api/v1/alerts.go @@ -236,7 +236,7 @@ func (h *AlertHandler) HandleCreateChannel(w http.ResponseWriter, r *http.Reques Headers string `json:"headers"` Secret string `json:"secret"` Config json.RawMessage `json:"config"` - Enabled bool `json:"enabled"` + Enabled *bool `json:"enabled"` } if err := json.NewDecoder(r.Body).Decode(&input); err != nil { WriteError(w, http.StatusBadRequest, "INVALID_BODY", "invalid JSON body") @@ -276,6 +276,11 @@ func (h *AlertHandler) HandleCreateChannel(w http.ResponseWriter, r *http.Reques } } + enabled := true + if input.Enabled != nil { + enabled = *input.Enabled + } + ch := &alert.NotificationChannel{ Name: input.Name, Type: input.Type, @@ -283,7 +288,7 @@ func (h *AlertHandler) HandleCreateChannel(w http.ResponseWriter, r *http.Reques Headers: input.Headers, Secret: input.Secret, Config: config, - Enabled: input.Enabled, + Enabled: enabled, } id, err := h.channelStore.InsertChannel(r.Context(), ch) diff --git a/internal/api/v1/alerts_test.go b/internal/api/v1/alerts_test.go index d79410ee..063cbcb2 100644 --- a/internal/api/v1/alerts_test.go +++ b/internal/api/v1/alerts_test.go @@ -190,6 +190,37 @@ func TestHandleCreateChannel_EmailValidation(t *testing.T) { } } +func TestHandleCreateChannel_EnabledDefaultsToTrue(t *testing.T) { + logger := slog.New(slog.NewTextHandler(io.Discard, nil)) + + cases := []struct { + name string + body string + wantEnabled bool + }{ + {"absent", `{"name":"hook","url":"https://example.com/hook"}`, true}, + {"explicit true", `{"name":"hook","url":"https://example.com/hook","enabled":true}`, true}, + {"explicit false", `{"name":"hook","url":"https://example.com/hook","enabled":false}`, false}, + } + + for _, tc := range cases { + t.Run(tc.name, func(t *testing.T) { + h := &AlertHandler{channelStore: &stubChannelStore{}, broker: NewSSEBroker(logger), allowPrivateWebhooks: true} + + req := httptest.NewRequest("POST", "/api/v1/channels", strings.NewReader(tc.body)) + req.Header.Set("Content-Type", "application/json") + rec := httptest.NewRecorder() + + h.HandleCreateChannel(rec, req) + + require.Equal(t, http.StatusCreated, rec.Code, rec.Body.String()) + var ch alert.NotificationChannel + require.NoError(t, json.Unmarshal(rec.Body.Bytes(), &ch)) + assert.Equal(t, tc.wantEnabled, ch.Enabled) + }) + } +} + // --------------------------------------------------------------------------- // HandleTestChannel — Pro channel type guard // --------------------------------------------------------------------------- diff --git a/internal/api/v1/containers.go b/internal/api/v1/containers.go index bde1d060..12462d4a 100644 --- a/internal/api/v1/containers.go +++ b/internal/api/v1/containers.go @@ -259,7 +259,6 @@ func (h *ContainerHandler) HandleGet(w http.ResponseWriter, r *http.Request) { "is_ignored": c.IsIgnored, "alert_severity": c.AlertSeverity, "restart_threshold": c.RestartThreshold, - "alert_channels": c.AlertChannels, "archived": c.Archived, "first_seen_at": c.FirstSeenAt, "last_state_change_at": c.LastStateChangeAt, diff --git a/internal/api/v1/escalation_test.go b/internal/api/v1/escalation_test.go index 4f578d9b..a2cdd259 100644 --- a/internal/api/v1/escalation_test.go +++ b/internal/api/v1/escalation_test.go @@ -199,7 +199,6 @@ func validPolicyBody() escalation.PolicyRequest { Filters: escalation.Filters{ Severities: []string{"critical"}, Scopes: []escalation.Scope{}, - Tags: []string{}, }, Levels: []escalation.LevelReq{ {DelaySeconds: 300, ChannelIDs: []string{"1"}}, @@ -406,7 +405,6 @@ func TestEscalation_UpdatePolicy_HappyPath(t *testing.T) { Filters: escalation.Filters{ Severities: []string{"critical"}, Scopes: []escalation.Scope{}, - Tags: []string{}, }, Levels: []escalation.LevelReq{ {DelaySeconds: 300, ChannelIDs: []string{"1"}}, diff --git a/internal/commercial/escalation/overlap.go b/internal/commercial/escalation/overlap.go index 1a237bf1..c6e6e655 100644 --- a/internal/commercial/escalation/overlap.go +++ b/internal/commercial/escalation/overlap.go @@ -44,8 +44,7 @@ func DetectOverlap(candidate *esc.Policy, existing []*esc.Policy) []esc.OverlapW // Empty list = universe (matches all). Non-empty list = explicit set. func filtersIntersect(a, b esc.Filters) bool { return setsIntersect(a.Severities, b.Severities) && - scopeSetsIntersect(a.Scopes, b.Scopes) && - setsIntersect(a.Tags, b.Tags) + scopeSetsIntersect(a.Scopes, b.Scopes) } // setsIntersect returns true if two string slices share at least one element, @@ -112,9 +111,6 @@ func filterIntersectionDescription(a, b esc.Filters) string { if len(a.Scopes) > 0 || len(b.Scopes) > 0 { parts = append(parts, "scopes") } - if len(a.Tags) > 0 || len(b.Tags) > 0 { - parts = append(parts, "tags") - } if len(parts) == 0 { return "all" } diff --git a/internal/commercial/escalation/overlap_test.go b/internal/commercial/escalation/overlap_test.go index 17004e12..90c71270 100644 --- a/internal/commercial/escalation/overlap_test.go +++ b/internal/commercial/escalation/overlap_test.go @@ -14,11 +14,11 @@ import ( func TestOverlap_BothEmpty_AllFilters_SharedChannel(t *testing.T) { a := &esc.Policy{ - Filters: esc.Filters{Severities: []string{}, Scopes: []esc.Scope{}, Tags: []string{}}, + Filters: esc.Filters{Severities: []string{}, Scopes: []esc.Scope{}}, Levels: []esc.Level{{ChannelIDs: []string{"1", "2"}}}, } b := &esc.Policy{ - Filters: esc.Filters{Severities: []string{}, Scopes: []esc.Scope{}, Tags: []string{}}, + Filters: esc.Filters{Severities: []string{}, Scopes: []esc.Scope{}}, Levels: []esc.Level{{ChannelIDs: []string{"2", "3"}}}, } warnings := DetectOverlap(a, []*esc.Policy{b}) @@ -28,11 +28,11 @@ func TestOverlap_BothEmpty_AllFilters_SharedChannel(t *testing.T) { func TestOverlap_NoSharedChannel(t *testing.T) { a := &esc.Policy{ - Filters: esc.Filters{Severities: []string{"critical"}, Scopes: []esc.Scope{}, Tags: []string{}}, + Filters: esc.Filters{Severities: []string{"critical"}, Scopes: []esc.Scope{}}, Levels: []esc.Level{{ChannelIDs: []string{"1"}}}, } b := &esc.Policy{ - Filters: esc.Filters{Severities: []string{"critical"}, Scopes: []esc.Scope{}, Tags: []string{}}, + Filters: esc.Filters{Severities: []string{"critical"}, Scopes: []esc.Scope{}}, Levels: []esc.Level{{ChannelIDs: []string{"2"}}}, } warnings := DetectOverlap(a, []*esc.Policy{b}) @@ -41,11 +41,11 @@ func TestOverlap_NoSharedChannel(t *testing.T) { func TestOverlap_DisjointSeverities(t *testing.T) { a := &esc.Policy{ - Filters: esc.Filters{Severities: []string{"warning"}, Scopes: []esc.Scope{}, Tags: []string{}}, + Filters: esc.Filters{Severities: []string{"warning"}, Scopes: []esc.Scope{}}, Levels: []esc.Level{{ChannelIDs: []string{"1"}}}, } b := &esc.Policy{ - Filters: esc.Filters{Severities: []string{"critical"}, Scopes: []esc.Scope{}, Tags: []string{}}, + Filters: esc.Filters{Severities: []string{"critical"}, Scopes: []esc.Scope{}}, Levels: []esc.Level{{ChannelIDs: []string{"1"}}}, } warnings := DetectOverlap(a, []*esc.Policy{b}) @@ -54,11 +54,11 @@ func TestOverlap_DisjointSeverities(t *testing.T) { func TestOverlap_OneEmptyFilters_IntersectsAll(t *testing.T) { a := &esc.Policy{ - Filters: esc.Filters{Severities: []string{}, Scopes: []esc.Scope{}, Tags: []string{}}, + Filters: esc.Filters{Severities: []string{}, Scopes: []esc.Scope{}}, Levels: []esc.Level{{ChannelIDs: []string{"1"}}}, } b := &esc.Policy{ - Filters: esc.Filters{Severities: []string{"critical"}, Scopes: []esc.Scope{}, Tags: []string{}}, + Filters: esc.Filters{Severities: []string{"critical"}, Scopes: []esc.Scope{}}, Levels: []esc.Level{{ChannelIDs: []string{"1"}}}, } warnings := DetectOverlap(a, []*esc.Policy{b}) @@ -68,12 +68,12 @@ func TestOverlap_OneEmptyFilters_IntersectsAll(t *testing.T) { func TestOverlap_SkipsSelf(t *testing.T) { a := &esc.Policy{ ID: "1", - Filters: esc.Filters{Severities: []string{}, Scopes: []esc.Scope{}, Tags: []string{}}, + Filters: esc.Filters{Severities: []string{}, Scopes: []esc.Scope{}}, Levels: []esc.Level{{ChannelIDs: []string{"1"}}}, } b := &esc.Policy{ ID: "1", - Filters: esc.Filters{Severities: []string{}, Scopes: []esc.Scope{}, Tags: []string{}}, + Filters: esc.Filters{Severities: []string{}, Scopes: []esc.Scope{}}, Levels: []esc.Level{{ChannelIDs: []string{"1"}}}, } warnings := DetectOverlap(a, []*esc.Policy{b}) @@ -82,14 +82,14 @@ func TestOverlap_SkipsSelf(t *testing.T) { func TestOverlap_MultiLevelSharedChannel(t *testing.T) { a := &esc.Policy{ - Filters: esc.Filters{Severities: []string{}, Scopes: []esc.Scope{}, Tags: []string{}}, + Filters: esc.Filters{Severities: []string{}, Scopes: []esc.Scope{}}, Levels: []esc.Level{ {ChannelIDs: []string{"10"}}, {ChannelIDs: []string{"5"}}, }, } b := &esc.Policy{ - Filters: esc.Filters{Severities: []string{}, Scopes: []esc.Scope{}, Tags: []string{}}, + Filters: esc.Filters{Severities: []string{}, Scopes: []esc.Scope{}}, Levels: []esc.Level{{ChannelIDs: []string{"5", "6"}}}, } warnings := DetectOverlap(a, []*esc.Policy{b}) diff --git a/internal/commercial/escalation/runner.go b/internal/commercial/escalation/runner.go index 77c4433e..267e1a2d 100644 --- a/internal/commercial/escalation/runner.go +++ b/internal/commercial/escalation/runner.go @@ -179,7 +179,7 @@ func (r *Runner) OnAlertCreated(ctx context.Context, a *alert.Alert) error { } // OnAlertAcknowledged terminates every active run attached to the alert and -// dispatches an ack notification on the channels of the last executed level. +// dispatches an ack notification on every channel the run has notified. func (r *Runner) OnAlertAcknowledged(ctx context.Context, alertID string, ack alert.Acknowledgment) error { runs, err := r.store.SelectActiveRunsByAlert(ctx, alertID) if err != nil { @@ -206,18 +206,39 @@ func (r *Runner) OnAlertAcknowledged(ctx context.Context, alertID string, ack al if a == nil || run.LastExecutedLevelIndex < 0 { continue } - policy, perr := unmarshalPolicySnapshot(run.PolicySnapshotJSON) - if perr != nil || policy == nil { + channelIDs, cErr := r.notifiedChannels(ctx, run.ID) + if cErr != nil { + r.logger.ErrorContext(ctx, "escalation: list notified channels for ack", "error", cErr, "run_id", run.ID) continue } - idx := run.LastExecutedLevelIndex - if idx >= len(policy.Levels) { + r.dispatchSpecial(ctx, run.ID, specialLevelAck, channelIDs, formatAckAlert(a, ack)) + } + return nil +} + +// notifiedChannels returns, once each and in delivery order, the channels a +// level of the run has sent to or is still sending to. +func (r *Runner) notifiedChannels(ctx context.Context, runID string) ([]string, error) { + deliveries, err := r.store.SelectRunDeliveries(ctx, runID) + if err != nil { + return nil, err + } + seen := make(map[string]struct{}, len(deliveries)) + var ids []string + for _, d := range deliveries { + if d.LevelIndex < 0 || d.ChannelID == nil { + continue + } + if d.Status != esc.DeliveryStatusSent && d.Status != esc.DeliveryStatusPending { continue } - ackAlert := formatAckAlert(a, ack) - r.dispatchSpecial(ctx, run.ID, specialLevelAck, policy.Levels[idx].ChannelIDs, ackAlert) + if _, dup := seen[*d.ChannelID]; dup { + continue + } + seen[*d.ChannelID] = struct{}{} + ids = append(ids, *d.ChannelID) } - return nil + return ids, nil } // OnAlertResolved terminates every active run attached to the alert. The @@ -326,12 +347,16 @@ func (r *Runner) processRun(ctx context.Context, run *esc.Run, now time.Time) er // Schedule the next tick. After the last level we still want the run to // surface once more so the next cycle can dispatch the "exhausted" notif - // and terminate the run — set next_action_at = now to make that happen. - var nextAt time.Time + // and terminate the run: next_action_at = now makes that happen. + nextAt := now if nextLevel+1 < len(policy.Levels) { - nextAt = now.Add(time.Duration(policy.Levels[nextLevel+1].DelaySeconds) * time.Second) - } else { - nextAt = now + // Step from the due slot, not from now: delays are cumulative and only a maintenance pause shifts them. + dueAt := now + if run.NextActionAt != nil { + dueAt = *run.NextActionAt + } + gap := policy.Levels[nextLevel+1].DelaySeconds - level.DelaySeconds + nextAt = dueAt.Add(time.Duration(gap) * time.Second) } if err := r.store.UpdateRunProgress(ctx, run.ID, nextLevel, &nextAt, esc.RunStatusActive); err != nil { return fmt.Errorf("update run progress: %w", err) @@ -455,8 +480,7 @@ func (r *Runner) abandonDelivery(ctx context.Context, d *esc.Delivery, reason st // --- helpers --- // matchPolicyFilters reports whether an alert satisfies a policy's filters. -// Empty filter buckets match everything (universe). Tags are not yet exposed -// on the Alert entity — treated as no-op (consistent with engine.matchesTrigger). +// Empty filter buckets match everything (universe). func matchPolicyFilters(a *alert.Alert, p *esc.Policy) bool { if len(p.Filters.Severities) > 0 && !slices.Contains(p.Filters.Severities, a.Severity) { return false @@ -493,24 +517,28 @@ func unmarshalPolicySnapshot(s string) (*esc.Policy, error) { func formatAckAlert(a *alert.Alert, ack alert.Acknowledgment) *alert.Alert { cp := *a cp.Status = alert.StatusResolved // routes the notifier through the "resolved" copy - at := ack.At.Format(time.RFC3339) - by := ack.By - if by == "" { - by = "—" + at := ack.At.UTC().Format(time.RFC3339) + if ack.By == "" { + cp.Message = fmt.Sprintf("Acknowledged at %s: %s", at, a.Message) + } else { + cp.Message = fmt.Sprintf("Acknowledged by %s at %s: %s", ack.By, at, a.Message) } - cp.Message = fmt.Sprintf("✓ Acquittée par %s à %s — %s", by, at, a.Message) return &cp } // formatExhaustedAlert produces a synthetic *Alert for the "escalation -// exhausted — human intervention required" notification (FR-013). +// exhausted, human intervention required" notification (FR-013). func formatExhaustedAlert(a *alert.Alert, totalLevels int) *alert.Alert { cp := *a // Keep severity/status of the original alert so the notifier still routes // it as a critical/warning event but with an explicit message. + levels := "levels" + if totalLevels == 1 { + levels = "level" + } cp.Message = fmt.Sprintf( - "⚠ Escalation épuisée — action humaine requise après %d palier(s) sans acquittement. %s", - totalLevels, a.Message, + "Escalation exhausted after %d %s without acknowledgment, human action required: %s", + totalLevels, levels, a.Message, ) return &cp } diff --git a/internal/commercial/escalation/runner_test.go b/internal/commercial/escalation/runner_test.go index ded01db2..02ae8839 100644 --- a/internal/commercial/escalation/runner_test.go +++ b/internal/commercial/escalation/runner_test.go @@ -9,6 +9,8 @@ import ( "errors" "io" "log/slog" + "slices" + "strings" "sync" "sync/atomic" "testing" @@ -112,8 +114,18 @@ func (s *runStore) CountActivePolicies(_ context.Context) (int, error) { return func (s *runStore) SelectRunsByPolicy(_ context.Context, _ string, _ int, _ string) ([]*esc.Run, error) { return nil, nil } -func (s *runStore) SelectRunDeliveries(_ context.Context, _ string) ([]*esc.Delivery, error) { - return nil, nil +func (s *runStore) SelectRunDeliveries(_ context.Context, runID string) ([]*esc.Delivery, error) { + s.mu.Lock() + defer s.mu.Unlock() + var out []*esc.Delivery + for _, d := range s.deliveries { + if d.RunID == runID { + cp := *d + out = append(out, &cp) + } + } + slices.SortFunc(out, func(a, b *esc.Delivery) int { return strings.Compare(a.ID, b.ID) }) + return out, nil } func (s *runStore) BulkDeactivateAllPolicies(_ context.Context) error { return nil } func (s *runStore) BulkRestorePoliciesFromDowngrade(_ context.Context) error { return nil } @@ -603,6 +615,7 @@ func TestRunner_EvaluateCycle_FiresLevelAndAdvances(t *testing.T) { a := criticalAlert("1") h.alerts.put(a) h.store.addPolicy(policyTwoLevels()) + start := h.now require.NoError(t, h.runner.OnAlertCreated(context.Background(), a)) h.advance(61 * time.Second) // first level due @@ -615,8 +628,7 @@ func TestRunner_EvaluateCycle_FiresLevelAndAdvances(t *testing.T) { assert.Equal(t, 0, r.LastExecutedLevelIndex) assert.Equal(t, esc.RunStatusActive, r.Status) require.NotNil(t, r.NextActionAt) - // next_action_at = current now + level[1].DelaySeconds (180s) - assert.Equal(t, h.now.Add(180*time.Second), *r.NextActionAt) + assert.Equal(t, start.Add(180*time.Second), *r.NextActionAt, "level delays count from the run start") calls := h.sender.snapshot() require.Len(t, calls, 1) @@ -694,6 +706,128 @@ func TestRunner_OnAlertAcknowledged_StopsRunsAndDispatchesAck(t *testing.T) { assert.Equal(t, 1, ackCount) } +func TestRunner_LevelsFireAtCumulativeDelaysDespiteLateTicks(t *testing.T) { + h := newHarness(t) + h.channels.put(defaultChannel()) + a := criticalAlert("1") + h.alerts.put(a) + h.store.addPolicy(&esc.Policy{ + Name: "p", Active: true, + Filters: esc.Filters{Severities: []string{alert.SeverityCritical}}, + Levels: []esc.Level{ + {Order: 0, DelaySeconds: 60, ChannelIDs: []string{"1"}}, + {Order: 1, DelaySeconds: 180, ChannelIDs: []string{"1"}}, + {Order: 2, DelaySeconds: 600, ChannelIDs: []string{"1"}}, + }, + }) + start := h.now + require.NoError(t, h.runner.OnAlertCreated(context.Background(), a)) + + h.advance(90 * time.Second) + require.NoError(t, h.runner.EvaluateCycle(context.Background())) + h.waitForDeliveries(t, 1) + runs := h.store.listRuns() + require.Len(t, runs, 1) + require.NotNil(t, runs[0].NextActionAt) + assert.Equal(t, start.Add(180*time.Second), *runs[0].NextActionAt) + + h.advance(110 * time.Second) + require.NoError(t, h.runner.EvaluateCycle(context.Background())) + h.waitForDeliveries(t, 2) + runs = h.store.listRuns() + assert.Equal(t, 1, runs[0].LastExecutedLevelIndex) + require.NotNil(t, runs[0].NextActionAt) + assert.Equal(t, start.Add(600*time.Second), *runs[0].NextActionAt) +} + +func TestRunner_MaintenancePauseShiftsLaterLevels(t *testing.T) { + h := newHarness(t) + h.channels.put(defaultChannel()) + a := criticalAlert("1") + h.alerts.put(a) + h.store.addPolicy(policyTwoLevels()) + start := h.now + require.NoError(t, h.runner.OnAlertCreated(context.Background(), a)) + + h.supp.set(true) + h.advance(61 * time.Second) + require.NoError(t, h.runner.EvaluateCycle(context.Background())) + + h.supp.set(false) + h.advance(239 * time.Second) + require.NoError(t, h.runner.EvaluateCycle(context.Background())) + require.NoError(t, h.runner.EvaluateCycle(context.Background())) + h.waitForDeliveries(t, 1) + + runs := h.store.listRuns() + require.Len(t, runs, 1) + assert.Equal(t, 0, runs[0].LastExecutedLevelIndex) + require.NotNil(t, runs[0].NextActionAt) + assert.Equal(t, start.Add(420*time.Second), *runs[0].NextActionAt, + "level 0 fired at start+300 once the window closed, level 1 keeps its 120s gap") +} + +func TestRunner_OnAlertAcknowledged_NotifiesEveryNotifiedChannelOnce(t *testing.T) { + h := newHarness(t) + for _, id := range []string{"1", "2", "3"} { + h.channels.put(&alert.NotificationChannel{ID: id, Name: "ch" + id, Type: "slack", URL: "u" + id, Enabled: true}) + } + h.channels.put(&alert.NotificationChannel{ID: "4", Name: "off", Type: "slack", URL: "u4", Enabled: false}) + a := criticalAlert("1") + h.alerts.put(a) + h.store.addPolicy(&esc.Policy{ + Name: "p", Active: true, + Filters: esc.Filters{Severities: []string{alert.SeverityCritical}}, + Levels: []esc.Level{ + {Order: 0, DelaySeconds: 60, ChannelIDs: []string{"1", "3", "4"}}, + {Order: 1, DelaySeconds: 180, ChannelIDs: []string{"1", "2"}}, + }, + }) + require.NoError(t, h.runner.OnAlertCreated(context.Background(), a)) + + h.advance(61 * time.Second) + require.NoError(t, h.runner.EvaluateCycle(context.Background())) + h.waitForDeliveries(t, 2) + h.advance(120 * time.Second) + require.NoError(t, h.runner.EvaluateCycle(context.Background())) + h.waitForDeliveries(t, 4) + + require.NoError(t, h.runner.OnAlertAcknowledged(context.Background(), "1", alert.Acknowledgment{By: "alice", At: h.now})) + h.waitForDeliveries(t, 7) + + var ackChannels []string + for _, d := range h.store.listDeliveries() { + if d.LevelIndex == specialLevelAck { + require.NotNil(t, d.ChannelID) + ackChannels = append(ackChannels, *d.ChannelID) + } + } + assert.ElementsMatch(t, []string{"1", "2", "3"}, ackChannels, + "every channel a level reached hears about the ack once; the disabled one never received anything") +} + +func TestFormatAckAlert_English(t *testing.T) { + a := criticalAlert("1") + at := time.Date(2026, 5, 7, 12, 30, 0, 0, time.UTC) + + assert.Equal(t, "Acknowledged by alice at 2026-05-07T12:30:00Z: container x stopped", + formatAckAlert(a, alert.Acknowledgment{By: "alice", At: at}).Message) + assert.Equal(t, "Acknowledged at 2026-05-07T12:30:00Z: container x stopped", + formatAckAlert(a, alert.Acknowledgment{At: at}).Message) + assert.Equal(t, "container x stopped", a.Message, "the source alert is left untouched") +} + +func TestFormatExhaustedAlert_English(t *testing.T) { + a := criticalAlert("1") + + assert.Equal(t, + "Escalation exhausted after 3 levels without acknowledgment, human action required: container x stopped", + formatExhaustedAlert(a, 3).Message) + assert.Equal(t, + "Escalation exhausted after 1 level without acknowledgment, human action required: container x stopped", + formatExhaustedAlert(a, 1).Message) +} + func TestRunner_OnAlertResolved_StopsRunsNoNotif(t *testing.T) { h := newHarness(t) h.channels.put(defaultChannel()) diff --git a/internal/commercial/escalation/service_test.go b/internal/commercial/escalation/service_test.go index 1668eb3d..5e4d54da 100644 --- a/internal/commercial/escalation/service_test.go +++ b/internal/commercial/escalation/service_test.go @@ -175,7 +175,6 @@ func validRequest() esc.PolicyRequest { Filters: esc.Filters{ Severities: []string{"critical"}, Scopes: []esc.Scope{}, - Tags: []string{}, }, Levels: []esc.LevelReq{ {DelaySeconds: 300, ChannelIDs: []string{"1"}}, @@ -353,7 +352,6 @@ func TestUpdatePolicy_HappyPath(t *testing.T) { Filters: esc.Filters{ Severities: []string{"warning"}, Scopes: []esc.Scope{}, - Tags: []string{}, }, Levels: []esc.LevelReq{ {DelaySeconds: 300, ChannelIDs: []string{"1"}}, diff --git a/internal/container/agent_event.go b/internal/container/agent_event.go index cb97df26..2d071a25 100644 --- a/internal/container/agent_event.go +++ b/internal/container/agent_event.go @@ -26,7 +26,6 @@ const ( labelPBGroup = "maintenant.group" labelPBSeverity = "maintenant.alert.severity" labelPBThreshold = "maintenant.alert.restart_threshold" - labelPBChannels = "maintenant.alert.channels" ) // HandleAgentEvent processes a ContainerEvent received from a remote agent. @@ -240,9 +239,6 @@ func applyAgentLabels(c *Container, labels map[string]string) { c.RestartThreshold = n } } - if v, ok := labels[labelPBChannels]; ok && v != "" { - c.AlertChannels = v - } } // agentEventTime returns the event's started_at when present, else now. diff --git a/internal/container/agent_event_test.go b/internal/container/agent_event_test.go index 32b6a912..82f0e814 100644 --- a/internal/container/agent_event_test.go +++ b/internal/container/agent_event_test.go @@ -247,7 +247,6 @@ func TestHandleAgentEvent_AppliesMaintenantLabels(t *testing.T) { labelPBGroup: "infra", labelPBSeverity: "critical", labelPBThreshold: "5", - labelPBChannels: "ops", }, } require.NoError(t, svc.HandleAgentEvent(context.Background(), "a", ev, agentevent.Meta{ObservedAt: time.Now()})) @@ -258,7 +257,6 @@ func TestHandleAgentEvent_AppliesMaintenantLabels(t *testing.T) { assert.Equal(t, "infra", c.CustomGroup) assert.Equal(t, SeverityCritical, c.AlertSeverity) assert.Equal(t, 5, c.RestartThreshold) - assert.Equal(t, "ops", c.AlertChannels) } func TestHandleAgentEvent_EmptyContainerIDIsNoOp(t *testing.T) { diff --git a/internal/container/model.go b/internal/container/model.go index f58837cb..2ae16f42 100644 --- a/internal/container/model.go +++ b/internal/container/model.go @@ -52,7 +52,6 @@ type Container struct { IsIgnored bool `json:"is_ignored"` AlertSeverity AlertSeverity `json:"alert_severity"` RestartThreshold int `json:"restart_threshold"` - AlertChannels string `json:"alert_channels,omitempty"` Archived bool `json:"archived"` FirstSeenAt time.Time `json:"first_seen_at"` LastStateChangeAt time.Time `json:"last_state_change_at"` diff --git a/internal/docker/discovery.go b/internal/docker/discovery.go index 615e00c8..b4c04ce7 100644 --- a/internal/docker/discovery.go +++ b/internal/docker/discovery.go @@ -26,7 +26,6 @@ const ( labelPBGroup = "maintenant.group" labelPBSeverity = "maintenant.alert.severity" labelPBThreshold = "maintenant.alert.restart_threshold" - labelPBChannels = "maintenant.alert.channels" ) // SecurityConfig holds security-relevant fields extracted from Docker's ContainerInspect. @@ -276,9 +275,6 @@ func applyLabels(cm *cmodel.Container, labels map[string]string) { cm.RestartThreshold = n } } - if v, ok := labels[labelPBChannels]; ok && v != "" { - cm.AlertChannels = v - } } func mapContainerState(state string) cmodel.ContainerState { diff --git a/internal/kubernetes/discovery.go b/internal/kubernetes/discovery.go index 87f75768..cae7ad3f 100644 --- a/internal/kubernetes/discovery.go +++ b/internal/kubernetes/discovery.go @@ -337,9 +337,6 @@ func applyAnnotations(cm *cmodel.Container, annotations map[string]string) { cm.RestartThreshold = n } } - if v, ok := annotations["maintenant.alert.channels"]; ok && v != "" { - cm.AlertChannels = v - } // Fallback display name from K8s standard labels. if cm.Name == "" { if v, ok := annotations["app.kubernetes.io/name"]; ok { diff --git a/internal/kubernetes/discovery_test.go b/internal/kubernetes/discovery_test.go index b5adf792..9e17900c 100644 --- a/internal/kubernetes/discovery_test.go +++ b/internal/kubernetes/discovery_test.go @@ -220,7 +220,6 @@ func TestDiscoverAll_Annotations(t *testing.T) { "maintenant.group": "backend", "maintenant.alert.severity": "critical", "maintenant.alert.restart_threshold": "5", - "maintenant.alert.channels": "slack", "maintenant.ignore": "true", }, }, @@ -260,9 +259,6 @@ func TestDiscoverAll_Annotations(t *testing.T) { if c.RestartThreshold != 5 { t.Errorf("expected RestartThreshold=5, got %d", c.RestartThreshold) } - if c.AlertChannels != "slack" { - t.Errorf("expected AlertChannels=slack, got %s", c.AlertChannels) - } if !c.IsIgnored { t.Error("expected IsIgnored=true") } diff --git a/internal/mcp/edition_conformance_test.go b/internal/mcp/edition_conformance_test.go index ad2348a3..a87eeaf3 100644 --- a/internal/mcp/edition_conformance_test.go +++ b/internal/mcp/edition_conformance_test.go @@ -1,9 +1,13 @@ package mcp import ( + "context" "encoding/json" + "regexp" + "strings" "testing" + gomcp "github.com/modelcontextprotocol/go-sdk/mcp" "github.com/stretchr/testify/assert" "github.com/stretchr/testify/require" @@ -57,6 +61,101 @@ func TestConformance_CheckCapabilityMatchesTheRegistry(t *testing.T) { } } +// A tool whose description says it requires an edition refuses below it. +func TestConformance_EveryGatedToolRefuses(t *testing.T) { + withEdition(t, extension.Community) + ctx := context.Background() + + server := newTestServer(t) + ct, st := gomcp.NewInMemoryTransports() + ss, err := server.Connect(ctx, st, nil) + require.NoError(t, err) + t.Cleanup(func() { _ = ss.Close() }) + cs, err := gomcp.NewClient(&gomcp.Implementation{Name: "conformance", Version: "0"}, nil).Connect(ctx, ct, nil) + require.NoError(t, err) + t.Cleanup(func() { _ = cs.Close() }) + + tools, err := cs.ListTools(ctx, nil) + require.NoError(t, err) + + gated := 0 + for _, tool := range tools.Tools { + m := requiresEdition.FindStringSubmatch(tool.Description) + if m == nil { + continue + } + required := extension.Edition(strings.ToLower(m[1])) + if !required.AtLeast(extension.Personal) { + continue + } + gated++ + t.Run(tool.Name, func(t *testing.T) { + result, err := cs.CallTool(ctx, &gomcp.CallToolParams{Name: tool.Name, Arguments: placeholderArgs(t, tool.InputSchema)}) + require.NoError(t, err) + require.True(t, result.IsError, "%s answered on Community", tool.Name) + + var payload struct { + Error string `json:"error"` + Feature string `json:"feature"` + RequiredEdition string `json:"required_edition"` + } + require.NoError(t, json.Unmarshal([]byte(textFromContent(t, result.Content)), &payload)) + assert.Equal(t, "edition_required", payload.Error) + assert.Equal(t, string(required), payload.RequiredEdition) + _, declared := extension.Catalog()[extension.Capability(payload.Feature)] + assert.True(t, declared, "feature %q is not a registered capability", payload.Feature) + }) + } + assert.NotZero(t, gated, "no gated tool found: the description convention changed") +} + +var requiresEdition = regexp.MustCompile(`Requires the (\w+) edition\.$`) + +// placeholderArgs fills each required property so the call gets past input validation. +func placeholderArgs(t *testing.T, schema any) map[string]any { + t.Helper() + raw, err := json.Marshal(schema) + require.NoError(t, err) + var s struct { + Required []string `json:"required"` + Properties map[string]struct { + Type any `json:"type"` + } `json:"properties"` + } + require.NoError(t, json.Unmarshal(raw, &s)) + + args := map[string]any{} + for _, name := range s.Required { + types := []string{} + switch v := s.Properties[name].Type.(type) { + case string: + types = append(types, v) + case []any: + for _, x := range v { + if str, ok := x.(string); ok && str != "null" { + types = append(types, str) + } + } + } + require.NotEmpty(t, types, "property %s has no type", name) + switch types[0] { + case "string": + args[name] = "x" + case "integer", "number": + args[name] = 1 + case "boolean": + args[name] = false + case "array": + args[name] = []any{} + case "object": + args[name] = map[string]any{} + default: + t.Fatalf("property %s has unexpected type %q", name, types[0]) + } + } + return args +} + // TestConformance_VocabularyIsTheRESTVocabulary pins the realignment. These // names are what makes the two surfaces comparable at all, so a drift here is // a contract break, not a cosmetic change. diff --git a/internal/mcp/tools_escalation.go b/internal/mcp/tools_escalation.go index 6230ea24..ac46fc2e 100644 --- a/internal/mcp/tools_escalation.go +++ b/internal/mcp/tools_escalation.go @@ -91,7 +91,6 @@ type createEscalationPolicyInput struct { type escalationFiltersInput struct { Severities []string `json:"severities,omitempty" jsonschema:"Severity filters: warning, critical"` Scopes []escalationScopeInput `json:"scopes,omitempty" jsonschema:"Scope filters"` - Tags []string `json:"tags,omitempty" jsonschema:"Tag filters"` } type escalationScopeInput struct { @@ -198,7 +197,6 @@ func createEscalationPolicyHandler(svc *Services) gomcp.ToolHandlerFor[createEsc Filters: escalation.Filters{ Severities: input.Filters.Severities, Scopes: scopes, - Tags: input.Filters.Tags, }, Levels: levels, } @@ -292,7 +290,6 @@ func updateEscalationPolicyHandler(svc *Services) gomcp.ToolHandlerFor[updateEsc Filters: escalation.Filters{ Severities: input.Filters.Severities, Scopes: scopes, - Tags: input.Filters.Tags, }, Levels: levels, } diff --git a/internal/mcp/tools_read.go b/internal/mcp/tools_read.go index 49957206..43745b4d 100644 --- a/internal/mcp/tools_read.go +++ b/internal/mcp/tools_read.go @@ -433,6 +433,9 @@ func getHealthHandler(svc *Services) gomcp.ToolHandlerFor[getHealthInput, any] { func listAgentsHandler(svc *Services) gomcp.ToolHandlerFor[listAgentsInput, any] { return func(ctx context.Context, _ *gomcp.CallToolRequest, _ listAgentsInput) (*gomcp.CallToolResult, any, error) { + if r, v, err := checkCapability(extension.CapMultihost); r != nil { + return r, v, err + } if svc.Agents == nil { return jsonResult([]any{}) } diff --git a/internal/mcp/tools_triggers.go b/internal/mcp/tools_triggers.go index 448d6c6a..5e0198c8 100644 --- a/internal/mcp/tools_triggers.go +++ b/internal/mcp/tools_triggers.go @@ -38,12 +38,12 @@ func registerTriggerTools(server *gomcp.Server, svc *Services) { addTool(server, svc, &gomcp.Tool{ Name: "create_trigger", - Description: "Create a new alert trigger, notifying recoveries by default. Scope and tag filters" + requires(extension.CapAlertAdvancedFilters), + Description: "Create a new alert trigger, notifying recoveries by default." + scopeFiltersRequire(), }, createTriggerHandler(svc)) addTool(server, svc, &gomcp.Tool{ Name: "update_trigger", - Description: "Update an existing alert trigger (last-write-wins). Scope and tag filters" + requires(extension.CapAlertAdvancedFilters), + Description: "Update an existing alert trigger (last-write-wins)." + scopeFiltersRequire(), }, updateTriggerHandler(svc)) addTool(server, svc, &gomcp.Tool{ @@ -65,7 +65,6 @@ type triggerInput struct { FilterSeverities string `json:"filter_severities" jsonschema:"CSV severity filter (e.g. 'critical,warning'). Empty matches everything."` FilterSources string `json:"filter_sources" jsonschema:"CSV source filter (e.g. 'container,endpoint'). Empty matches everything."` FilterScopes string `json:"filter_scopes" jsonschema:"CSV scope filter (e.g. 'container:42,endpoint:7'). Needs the advanced filters capability. Empty matches everything."` - FilterTags string `json:"filter_tags" jsonschema:"CSV tag filter. Needs the advanced filters capability. Empty matches everything."` Enabled bool `json:"enabled" jsonschema:"Whether the trigger is active"` NotifyOnResolve *bool `json:"notify_on_resolve,omitempty" jsonschema:"Relay recovery (resolved) notifications, default true"` ChannelIDs []string `json:"channel_ids" jsonschema:"Notification channel IDs (at least one required)"` @@ -82,11 +81,16 @@ type deleteTriggerInput struct { // --- helpers --- -// checkAdvancedFilters refuses scope/tag filters the running edition does not +// scopeFiltersRequire tells, in a tool description, which edition scope filters need. +func scopeFiltersRequire() string { + return " Scope filters require the " + titleEdition(extension.MinEdition(extension.CapAlertAdvancedFilters)) + " edition." +} + +// checkAdvancedFilters refuses scope filters the running edition does not // open. Only the "are advanced filters even in play" shortcut lives here; the // edition decision itself goes through the registry like every other. -func checkAdvancedFilters(scopes, tags string) (*gomcp.CallToolResult, any, error) { - if scopes == "" && tags == "" { +func checkAdvancedFilters(scopes string) (*gomcp.CallToolResult, any, error) { + if scopes == "" { return nil, nil, nil } return checkCapability(extension.CapAlertAdvancedFilters) @@ -144,7 +148,7 @@ func createTriggerHandler(svc *Services) gomcp.ToolHandlerFor[triggerInput, any] if svc.Triggers == nil || svc.Channels == nil { return errResult("trigger or channel store not available") } - if r, v, err := checkAdvancedFilters(input.FilterScopes, input.FilterTags); r != nil { + if r, v, err := checkAdvancedFilters(input.FilterScopes); r != nil { return r, v, err } if err := validateTriggerCommon(&input); err != nil { @@ -170,7 +174,6 @@ func createTriggerHandler(svc *Services) gomcp.ToolHandlerFor[triggerInput, any] FilterSeverities: input.FilterSeverities, FilterSources: input.FilterSources, FilterScopes: input.FilterScopes, - FilterTags: input.FilterTags, Enabled: input.Enabled, NotifyOnResolve: notifyOnResolve, ChannelIDs: input.ChannelIDs, @@ -198,7 +201,7 @@ func updateTriggerHandler(svc *Services) gomcp.ToolHandlerFor[updateTriggerInput if existing == nil { return errResult("trigger not found") } - if r, v, err := checkAdvancedFilters(input.FilterScopes, input.FilterTags); r != nil { + if r, v, err := checkAdvancedFilters(input.FilterScopes); r != nil { return r, v, err } if err := validateTriggerCommon(&input.triggerInput); err != nil { @@ -218,7 +221,6 @@ func updateTriggerHandler(svc *Services) gomcp.ToolHandlerFor[updateTriggerInput existing.FilterSeverities = input.FilterSeverities existing.FilterSources = input.FilterSources existing.FilterScopes = input.FilterScopes - existing.FilterTags = input.FilterTags existing.Enabled = input.Enabled if input.NotifyOnResolve != nil { existing.NotifyOnResolve = *input.NotifyOnResolve diff --git a/internal/mcp/tools_triggers_test.go b/internal/mcp/tools_triggers_test.go index 44ac88b0..10ce80e7 100644 --- a/internal/mcp/tools_triggers_test.go +++ b/internal/mcp/tools_triggers_test.go @@ -7,6 +7,7 @@ import ( "context" "encoding/json" "strconv" + "strings" "testing" "time" @@ -290,22 +291,21 @@ func TestCreateTriggerHandler_FilterScopes_CommunityBlocked(t *testing.T) { assert.Contains(t, textFromContent(t, result.Content), "edition_required") } -func TestCreateTriggerHandler_FilterTags_CommunityBlocked(t *testing.T) { - original := extension.CurrentEdition - extension.CurrentEdition = func() extension.Edition { return extension.Community } - defer func() { extension.CurrentEdition = original }() - - svc, _ := buildTriggerServices() - handler := createTriggerHandler(svc) - - result, _, err := handler(context.Background(), nil, triggerInput{ - Name: "TaggedTrigger", - FilterTags: "prod", - ChannelIDs: []string{"1"}, - }) - require.NoError(t, err) - require.True(t, result.IsError) - assert.Contains(t, textFromContent(t, result.Content), "edition_required") +func TestFilterTools_HaveNoTagFilter(t *testing.T) { + tools := listToolNames(t, newTestServer(t)) + for name, property := range map[string]string{ + "create_trigger": `"filter_tags"`, + "update_trigger": `"filter_tags"`, + "create_escalation_policy": `"tags"`, + "update_escalation_policy": `"tags"`, + } { + tool := tools[name] + require.NotNil(t, tool, name) + schema, err := json.Marshal(tool.InputSchema) + require.NoError(t, err) + assert.NotContains(t, string(schema), property, "%s must not offer a tag filter", name) + assert.NotContains(t, strings.ToLower(tool.Description), "tag", name) + } } func TestCreateTriggerHandler_FilterScopes_ProAllowed(t *testing.T) { diff --git a/internal/store/containers.go b/internal/store/containers.go index 8ff01b3c..49b9fbaa 100644 --- a/internal/store/containers.go +++ b/internal/store/containers.go @@ -37,18 +37,18 @@ func (s *ContainerStore) InsertContainer(ctx context.Context, c *container.Conta _, err := s.writer.Exec(ctx, `INSERT INTO containers (id, agent_id, external_id, name, image, state, health_status, has_health_check, orchestration_group, orchestration_unit, custom_group, is_ignored, alert_severity, - restart_threshold, alert_channels, archived, first_seen_at, last_state_change_at, + restart_threshold, archived, first_seen_at, last_state_change_at, runtime_type, error_detail, controller_kind, namespace, pod_count, ready_count, compose_working_dir, swarm_service_id, swarm_service_name, swarm_service_mode, swarm_node_id, swarm_task_slot, swarm_desired_replicas, image_version, image_source, image_url, image_description) - VALUES (?, ?, ?, ?, ?, ?, ?, ?, ?, ?, ?, ?, ?, ?, ?, ?, ?, ?, ?, ?, ?, ?, ?, ?, ?, ?, ?, ?, ?, ?, ?, ?, ?, ?, ?) + VALUES (?, ?, ?, ?, ?, ?, ?, ?, ?, ?, ?, ?, ?, ?, ?, ?, ?, ?, ?, ?, ?, ?, ?, ?, ?, ?, ?, ?, ?, ?, ?, ?, ?, ?) ON CONFLICT(id) DO UPDATE SET name=excluded.name, image=excluded.image, state=excluded.state, health_status=excluded.health_status, has_health_check=excluded.has_health_check, orchestration_group=excluded.orchestration_group, orchestration_unit=excluded.orchestration_unit, custom_group=excluded.custom_group, is_ignored=excluded.is_ignored, alert_severity=excluded.alert_severity, - restart_threshold=excluded.restart_threshold, alert_channels=excluded.alert_channels, + restart_threshold=excluded.restart_threshold, archived=excluded.archived, last_state_change_at=excluded.last_state_change_at, agent_id=excluded.agent_id, runtime_type=excluded.runtime_type, error_detail=excluded.error_detail, @@ -63,7 +63,7 @@ func (s *ContainerStore) InsertContainer(ctx context.Context, c *container.Conta c.ID, c.AgentID, c.ExternalID, c.Name, c.Image, string(c.State), nullableHealth(c.HealthStatus), boolToInt(c.HasHealthCheck), NullableString(c.OrchestrationGroup), NullableString(c.OrchestrationUnit), NullableString(c.CustomGroup), boolToInt(c.IsIgnored), string(c.AlertSeverity), - c.RestartThreshold, NullableString(c.AlertChannels), boolToInt(c.Archived), + c.RestartThreshold, boolToInt(c.Archived), c.FirstSeenAt.Unix(), c.LastStateChangeAt.Unix(), c.RuntimeType, c.ErrorDetail, c.ControllerKind, c.Namespace, c.PodCount, c.ReadyCount, c.ComposeWorkingDir, @@ -81,7 +81,7 @@ func (s *ContainerStore) UpdateContainer(ctx context.Context, c *container.Conta _, err := s.writer.Exec(ctx, `UPDATE containers SET name=?, image=?, state=?, health_status=?, has_health_check=?, orchestration_group=?, orchestration_unit=?, custom_group=?, is_ignored=?, alert_severity=?, - restart_threshold=?, alert_channels=?, archived=?, last_state_change_at=?, archived_at=?, + restart_threshold=?, archived=?, last_state_change_at=?, archived_at=?, runtime_type=?, error_detail=?, controller_kind=?, namespace=?, pod_count=?, ready_count=?, compose_working_dir=?, swarm_service_id=?, swarm_service_name=?, swarm_service_mode=?, swarm_node_id=?, swarm_task_slot=?, swarm_desired_replicas=?, @@ -91,7 +91,7 @@ func (s *ContainerStore) UpdateContainer(ctx context.Context, c *container.Conta c.Name, c.Image, string(c.State), nullableHealth(c.HealthStatus), boolToInt(c.HasHealthCheck), NullableString(c.OrchestrationGroup), NullableString(c.OrchestrationUnit), NullableString(c.CustomGroup), boolToInt(c.IsIgnored), string(c.AlertSeverity), - c.RestartThreshold, NullableString(c.AlertChannels), boolToInt(c.Archived), + c.RestartThreshold, boolToInt(c.Archived), c.LastStateChangeAt.Unix(), nullableTime(c.ArchivedAt), c.RuntimeType, c.ErrorDetail, c.ControllerKind, c.Namespace, c.PodCount, c.ReadyCount, c.ComposeWorkingDir, @@ -375,7 +375,7 @@ func (s *ContainerStore) DeleteArchivedContainersBefore(ctx context.Context, bef const containerColumns = `id, agent_id, external_id, name, image, state, health_status, has_health_check, orchestration_group, orchestration_unit, custom_group, is_ignored, alert_severity, - restart_threshold, alert_channels, archived, first_seen_at, last_state_change_at, archived_at, + restart_threshold, archived, first_seen_at, last_state_change_at, archived_at, runtime_type, error_detail, controller_kind, namespace, pod_count, ready_count, compose_working_dir, swarm_service_id, swarm_service_name, swarm_service_mode, swarm_node_id, swarm_task_slot, swarm_desired_replicas, @@ -389,7 +389,7 @@ type rowScanner interface { func (s *ContainerStore) scanContainer(row rowScanner) (*container.Container, error) { var c container.Container - var healthStatus, orchestrationGroup, orchestrationUnit, customGroup, alertChannels sql.NullString + var healthStatus, orchestrationGroup, orchestrationUnit, customGroup sql.NullString var hasHealthCheck, isIgnored, archived int var firstSeen, lastChange int64 var archivedAt sql.NullInt64 @@ -399,7 +399,7 @@ func (s *ContainerStore) scanContainer(row rowScanner) (*container.Container, er &healthStatus, &hasHealthCheck, &orchestrationGroup, &orchestrationUnit, &customGroup, &isIgnored, &c.AlertSeverity, - &c.RestartThreshold, &alertChannels, + &c.RestartThreshold, &archived, &firstSeen, &lastChange, &archivedAt, &c.RuntimeType, &c.ErrorDetail, &c.ControllerKind, &c.Namespace, &c.PodCount, &c.ReadyCount, &c.ComposeWorkingDir, @@ -432,9 +432,6 @@ func (s *ContainerStore) scanContainer(row rowScanner) (*container.Container, er if customGroup.Valid { c.CustomGroup = customGroup.String } - if alertChannels.Valid { - c.AlertChannels = alertChannels.String - } if archivedAt.Valid { t := time.Unix(archivedAt.Int64, 0) c.ArchivedAt = &t diff --git a/internal/store/escalation.go b/internal/store/escalation.go index 05177bd9..92ab391f 100644 --- a/internal/store/escalation.go +++ b/internal/store/escalation.go @@ -38,10 +38,6 @@ func (s *EscalationStore) InsertPolicy(ctx context.Context, p *escalation.Policy if err != nil { return "", fmt.Errorf("marshal scopes: %w", err) } - tagsJSON, err := json.Marshal(p.Filters.Tags) - if err != nil { - return "", fmt.Errorf("marshal tags: %w", err) - } levelsJSON, err := json.Marshal(p.Levels) if err != nil { return "", fmt.Errorf("marshal levels: %w", err) @@ -53,10 +49,10 @@ func (s *EscalationStore) InsertPolicy(ctx context.Context, p *escalation.Policy } _, err = s.writer.Exec(ctx, `INSERT INTO escalation_policies - (id, name, active, active_before_downgrade, severities_json, scopes_json, tags_json, levels_json, created_at, created_by, updated_at, updated_by) - VALUES (?, ?, ?, ?, ?, ?, ?, ?, ?, ?, ?, ?)`, + (id, name, active, active_before_downgrade, severities_json, scopes_json, levels_json, created_at, created_by, updated_at, updated_by) + VALUES (?, ?, ?, ?, ?, ?, ?, ?, ?, ?, ?)`, p.ID, p.Name, boolToInt(p.Active), boolToInt(p.ActiveBeforeDowngrade), - string(sevJSON), string(scopesJSON), string(tagsJSON), string(levelsJSON), + string(sevJSON), string(scopesJSON), string(levelsJSON), p.CreatedAt.Unix(), NullableString(p.CreatedBy), p.CreatedAt.Unix(), NullableString(p.UpdatedBy), ) if err != nil { @@ -74,10 +70,6 @@ func (s *EscalationStore) UpdatePolicy(ctx context.Context, p *escalation.Policy if err != nil { return fmt.Errorf("marshal scopes: %w", err) } - tagsJSON, err := json.Marshal(p.Filters.Tags) - if err != nil { - return fmt.Errorf("marshal tags: %w", err) - } levelsJSON, err := json.Marshal(p.Levels) if err != nil { return fmt.Errorf("marshal levels: %w", err) @@ -86,11 +78,11 @@ func (s *EscalationStore) UpdatePolicy(ctx context.Context, p *escalation.Policy _, err = s.writer.Exec(ctx, `UPDATE escalation_policies SET name=?, active=?, active_before_downgrade=?, - severities_json=?, scopes_json=?, tags_json=?, levels_json=?, + severities_json=?, scopes_json=?, levels_json=?, updated_at=?, updated_by=? WHERE id=?`, p.Name, boolToInt(p.Active), boolToInt(p.ActiveBeforeDowngrade), - string(sevJSON), string(scopesJSON), string(tagsJSON), string(levelsJSON), + string(sevJSON), string(scopesJSON), string(levelsJSON), time.Now().Unix(), NullableString(p.UpdatedBy), p.ID, ) @@ -103,7 +95,7 @@ func (s *EscalationStore) UpdatePolicy(ctx context.Context, p *escalation.Policy func (s *EscalationStore) SelectPolicy(ctx context.Context, id string) (*escalation.Policy, error) { row := s.db.QueryRowContext(ctx, `SELECT id, name, active, active_before_downgrade, - severities_json, scopes_json, tags_json, levels_json, + severities_json, scopes_json, levels_json, created_at, created_by, updated_at, updated_by FROM escalation_policies WHERE id = ?`, id) p, err := scanEscalationPolicy(row) @@ -115,7 +107,7 @@ func (s *EscalationStore) SelectPolicy(ctx context.Context, id string) (*escalat func (s *EscalationStore) SelectPolicies(ctx context.Context, activeOnly bool) ([]*escalation.Policy, error) { q := `SELECT id, name, active, active_before_downgrade, - severities_json, scopes_json, tags_json, levels_json, + severities_json, scopes_json, levels_json, created_at, created_by, updated_at, updated_by FROM escalation_policies` if activeOnly { @@ -481,13 +473,13 @@ func (s *EscalationStore) PurgeRunsAndDeliveriesOlderThan(ctx context.Context, b func scanEscalationPolicy(scanner rowScanner) (*escalation.Policy, error) { var p escalation.Policy var active, activeBeforeDowngrade int - var sevJSON, scopesJSON, tagsJSON, levelsJSON string + var sevJSON, scopesJSON, levelsJSON string var createdAt, updatedAt int64 var createdBy, updatedBy sql.NullString err := scanner.Scan( &p.ID, &p.Name, &active, &activeBeforeDowngrade, - &sevJSON, &scopesJSON, &tagsJSON, &levelsJSON, + &sevJSON, &scopesJSON, &levelsJSON, &createdAt, &createdBy, &updatedAt, &updatedBy, ) if err != nil { @@ -507,7 +499,6 @@ func scanEscalationPolicy(scanner rowScanner) (*escalation.Policy, error) { _ = json.Unmarshal([]byte(sevJSON), &p.Filters.Severities) _ = json.Unmarshal([]byte(scopesJSON), &p.Filters.Scopes) - _ = json.Unmarshal([]byte(tagsJSON), &p.Filters.Tags) _ = json.Unmarshal([]byte(levelsJSON), &p.Levels) if p.Filters.Severities == nil { @@ -516,9 +507,6 @@ func scanEscalationPolicy(scanner rowScanner) (*escalation.Policy, error) { if p.Filters.Scopes == nil { p.Filters.Scopes = []escalation.Scope{} } - if p.Filters.Tags == nil { - p.Filters.Tags = []string{} - } if p.Levels == nil { p.Levels = []escalation.Level{} } diff --git a/internal/store/escalation_test.go b/internal/store/escalation_test.go index f2769a24..71469de1 100644 --- a/internal/store/escalation_test.go +++ b/internal/store/escalation_test.go @@ -64,7 +64,6 @@ func setupEscalationTestDB(t *testing.T) (*EscalationStore, *sql.DB) { active_before_downgrade INTEGER NOT NULL DEFAULT 0, severities_json TEXT NOT NULL DEFAULT '[]', scopes_json TEXT NOT NULL DEFAULT '[]', - tags_json TEXT NOT NULL DEFAULT '[]', levels_json TEXT NOT NULL, created_at BIGINT NOT NULL DEFAULT 0, created_by TEXT, @@ -119,7 +118,6 @@ func makeTestPolicy(name string, active bool) *escalation.Policy { Filters: escalation.Filters{ Severities: []string{"critical"}, Scopes: []escalation.Scope{{Kind: "container", RefID: "1"}}, - Tags: []string{"prod"}, }, Levels: []escalation.Level{ {Order: 0, DelaySeconds: 300, ChannelIDs: []string{"1", "2"}}, diff --git a/internal/store/migrations/postgres/36_drop_tag_filters_and_alert_channels.down.sql b/internal/store/migrations/postgres/36_drop_tag_filters_and_alert_channels.down.sql new file mode 100644 index 00000000..3ccc038b --- /dev/null +++ b/internal/store/migrations/postgres/36_drop_tag_filters_and_alert_channels.down.sql @@ -0,0 +1,3 @@ +ALTER TABLE containers ADD COLUMN alert_channels TEXT; +ALTER TABLE escalation_policies ADD COLUMN tags_json TEXT NOT NULL DEFAULT '[]'; +ALTER TABLE alert_triggers ADD COLUMN filter_tags TEXT NOT NULL DEFAULT ''; diff --git a/internal/store/migrations/postgres/36_drop_tag_filters_and_alert_channels.up.sql b/internal/store/migrations/postgres/36_drop_tag_filters_and_alert_channels.up.sql new file mode 100644 index 00000000..65829526 --- /dev/null +++ b/internal/store/migrations/postgres/36_drop_tag_filters_and_alert_channels.up.sql @@ -0,0 +1,3 @@ +ALTER TABLE alert_triggers DROP COLUMN filter_tags; +ALTER TABLE escalation_policies DROP COLUMN tags_json; +ALTER TABLE containers DROP COLUMN alert_channels; diff --git a/internal/store/migrations/sqlite/36_drop_tag_filters_and_alert_channels.down.sql b/internal/store/migrations/sqlite/36_drop_tag_filters_and_alert_channels.down.sql new file mode 100644 index 00000000..3ccc038b --- /dev/null +++ b/internal/store/migrations/sqlite/36_drop_tag_filters_and_alert_channels.down.sql @@ -0,0 +1,3 @@ +ALTER TABLE containers ADD COLUMN alert_channels TEXT; +ALTER TABLE escalation_policies ADD COLUMN tags_json TEXT NOT NULL DEFAULT '[]'; +ALTER TABLE alert_triggers ADD COLUMN filter_tags TEXT NOT NULL DEFAULT ''; diff --git a/internal/store/migrations/sqlite/36_drop_tag_filters_and_alert_channels.up.sql b/internal/store/migrations/sqlite/36_drop_tag_filters_and_alert_channels.up.sql new file mode 100644 index 00000000..65829526 --- /dev/null +++ b/internal/store/migrations/sqlite/36_drop_tag_filters_and_alert_channels.up.sql @@ -0,0 +1,3 @@ +ALTER TABLE alert_triggers DROP COLUMN filter_tags; +ALTER TABLE escalation_policies DROP COLUMN tags_json; +ALTER TABLE containers DROP COLUMN alert_channels; diff --git a/internal/store/migrations_postgres_test.go b/internal/store/migrations_postgres_test.go index 024a632c..7abf8d08 100644 --- a/internal/store/migrations_postgres_test.go +++ b/internal/store/migrations_postgres_test.go @@ -111,7 +111,10 @@ func TestMigratePostgres_ConcurrentCatchUp(t *testing.T) { "DROP TABLE cve_evaluations", // 31 "ALTER TABLE containers DROP COLUMN image_version, DROP COLUMN image_source, DROP COLUMN image_url, DROP COLUMN image_description", // 32 "ALTER TABLE agents DROP COLUMN os_id, DROP COLUMN os_version_id, DROP COLUMN os_pretty_name, DROP COLUMN os_source, DROP COLUMN os_unavailable_reason, DROP COLUMN os_reported_at", // 33 - "DROP TABLE outbound_heartbeats", // 34 + "DROP TABLE outbound_heartbeats", // 34 + "ALTER TABLE containers ADD COLUMN alert_channels TEXT", // 36 + "ALTER TABLE escalation_policies ADD COLUMN tags_json TEXT NOT NULL DEFAULT '[]'", // 36 + "ALTER TABLE alert_triggers ADD COLUMN filter_tags TEXT NOT NULL DEFAULT ''", // 36 } { _, err = db.ReadDB().Exec(undo) require.NoError(t, err, undo) diff --git a/internal/store/migrations_test.go b/internal/store/migrations_test.go index dfa96592..040fb6db 100644 --- a/internal/store/migrations_test.go +++ b/internal/store/migrations_test.go @@ -6,6 +6,7 @@ package store import ( "context" "database/sql" + "io/fs" "os" "path/filepath" "testing" @@ -212,3 +213,44 @@ func TestMigrateSQLite_DirtyFirstMigrationRecovers(t *testing.T) { require.NoError(t, err) assert.Equal(t, head, v, "the interrupted install must be replayed to the head") } + +func TestMigration36_DropsTagFiltersAndAlertChannels(t *testing.T) { + db := openTestDB(t) + ctx := context.Background() + dropped := map[string]string{ + "alert_triggers": "filter_tags", + "escalation_policies": "tags_json", + "containers": "alert_channels", + } + + assertColumns := func(present bool) { + t.Helper() + desc := introspectSQLite + if db.dialect == DialectPostgres { + desc = introspectPostgres + } + tables := desc(t, db.ReadDB()).tables + for table, col := range dropped { + require.Contains(t, tables, table) + if present { + assert.Contains(t, tables[table], col) + } else { + assert.NotContains(t, tables[table], col) + } + } + } + apply := func(direction string) { + t.Helper() + sqlText, err := fs.ReadFile(migrationFS, + "migrations/"+db.dialect.String()+"/36_drop_tag_filters_and_alert_channels."+direction+".sql") + require.NoError(t, err) + _, err = db.ReadDB().ExecContext(ctx, string(sqlText)) + require.NoError(t, err, direction) + } + + assertColumns(false) + apply("down") + assertColumns(true) + apply("up") + assertColumns(false) +} diff --git a/internal/store/transform.go b/internal/store/transform.go index 3e8f8d32..883f8880 100644 --- a/internal/store/transform.go +++ b/internal/store/transform.go @@ -239,13 +239,13 @@ func copyStatements() []stmt { {"containers", `INSERT INTO containers (id, agent_id, external_id, name, image, state, health_status, has_health_check, orchestration_group, orchestration_unit, custom_group, is_ignored, alert_severity, - restart_threshold, alert_channels, archived, first_seen_at, last_state_change_at, archived_at, + restart_threshold, archived, first_seen_at, last_state_change_at, archived_at, runtime_type, error_detail, controller_kind, namespace, pod_count, ready_count, compose_working_dir, swarm_service_id, swarm_service_name, swarm_service_mode, swarm_node_id, swarm_task_slot, swarm_desired_replicas) SELECT mnt_container_id('` + s + `', external_id), '` + s + `', external_id, name, image, state, health_status, has_health_check, orchestration_group, orchestration_unit, custom_group, - is_ignored, alert_severity, restart_threshold, alert_channels, archived, first_seen_at, + is_ignored, alert_severity, restart_threshold, archived, first_seen_at, last_state_change_at, archived_at, runtime_type, error_detail, controller_kind, namespace, pod_count, ready_count, compose_working_dir, swarm_service_id, swarm_service_name, swarm_service_mode, swarm_node_id, swarm_task_slot, swarm_desired_replicas @@ -418,8 +418,8 @@ func copyStatements() []stmt { // -------- alert triggers (minted) + channels join ----------------------- {"alert_triggers", `INSERT INTO alert_triggers - (id, name, filter_severities, filter_sources, filter_scopes, filter_tags, enabled, notify_on_resolve, created_at, updated_at) - SELECT mt.new_id, t.name, t.filter_severities, t.filter_sources, t.filter_scopes, t.filter_tags, + (id, name, filter_severities, filter_sources, filter_scopes, enabled, notify_on_resolve, created_at, updated_at) + SELECT mt.new_id, t.name, t.filter_severities, t.filter_sources, t.filter_scopes, t.enabled, t.notify_on_resolve, ` + epoch("t.created_at") + `, ` + epoch("t.updated_at") + ` FROM _old_alert_triggers t JOIN _map_trigger mt ON t.id = mt.old_id`}, @@ -430,10 +430,10 @@ func copyStatements() []stmt { // -------- escalation policies/runs/deliveries (minted) ------------------ {"escalation_policies", `INSERT INTO escalation_policies - (id, name, active, active_before_downgrade, severities_json, scopes_json, tags_json, levels_json, + (id, name, active, active_before_downgrade, severities_json, scopes_json, levels_json, created_at, created_by, updated_at, updated_by) SELECT mp.new_id, p.name, p.active, p.active_before_downgrade, p.severities_json, p.scopes_json, - p.tags_json, p.levels_json, ` + epoch("p.created_at") + `, p.created_by, + p.levels_json, ` + epoch("p.created_at") + `, p.created_by, ` + epoch("p.updated_at") + `, p.updated_by FROM _old_escalation_policies p JOIN _map_policy mp ON p.id = mp.old_id`}, diff --git a/internal/store/triggers.go b/internal/store/triggers.go index e3f72d47..d14af5fc 100644 --- a/internal/store/triggers.go +++ b/internal/store/triggers.go @@ -34,9 +34,9 @@ func (s *TriggerStoreImpl) InsertTrigger(ctx context.Context, t *alert.AlertTrig now := time.Now().Unix() _, err := s.writer.Exec(ctx, `INSERT INTO alert_triggers - (id, name, filter_severities, filter_sources, filter_scopes, filter_tags, enabled, notify_on_resolve, created_at, updated_at) - VALUES (?, ?, ?, ?, ?, ?, ?, ?, ?, ?)`, - t.ID, t.Name, t.FilterSeverities, t.FilterSources, t.FilterScopes, t.FilterTags, boolToInt(t.Enabled), boolToInt(t.NotifyOnResolve), now, now, + (id, name, filter_severities, filter_sources, filter_scopes, enabled, notify_on_resolve, created_at, updated_at) + VALUES (?, ?, ?, ?, ?, ?, ?, ?, ?)`, + t.ID, t.Name, t.FilterSeverities, t.FilterSources, t.FilterScopes, boolToInt(t.Enabled), boolToInt(t.NotifyOnResolve), now, now, ) if err != nil { return "", fmt.Errorf("insert trigger: %w", err) @@ -49,7 +49,7 @@ func (s *TriggerStoreImpl) InsertTrigger(ctx context.Context, t *alert.AlertTrig func (s *TriggerStoreImpl) GetTrigger(ctx context.Context, id string) (*alert.AlertTrigger, error) { row := s.db.QueryRowContext(ctx, - `SELECT id, name, filter_severities, filter_sources, filter_scopes, filter_tags, + `SELECT id, name, filter_severities, filter_sources, filter_scopes, enabled, notify_on_resolve, created_at, updated_at FROM alert_triggers WHERE id = ?`, id) @@ -80,7 +80,7 @@ func (s *TriggerStoreImpl) ListEnabledTriggers(ctx context.Context) ([]*alert.Al func (s *TriggerStoreImpl) listTriggersWhere(ctx context.Context, where string) ([]*alert.AlertTrigger, error) { // #nosec G202 -- `where` is a package-internal constant, never caller input. - query := `SELECT id, name, filter_severities, filter_sources, filter_scopes, filter_tags, + query := `SELECT id, name, filter_severities, filter_sources, filter_scopes, enabled, notify_on_resolve, created_at, updated_at FROM alert_triggers ` + where + ` ORDER BY created_at ASC` @@ -117,10 +117,10 @@ func (s *TriggerStoreImpl) listTriggersWhere(ctx context.Context, where string) func (s *TriggerStoreImpl) UpdateTrigger(ctx context.Context, t *alert.AlertTrigger) error { _, err := s.writer.Exec(ctx, `UPDATE alert_triggers - SET name=?, filter_severities=?, filter_sources=?, filter_scopes=?, filter_tags=?, + SET name=?, filter_severities=?, filter_sources=?, filter_scopes=?, enabled=?, notify_on_resolve=?, updated_at=? WHERE id=?`, - t.Name, t.FilterSeverities, t.FilterSources, t.FilterScopes, t.FilterTags, + t.Name, t.FilterSeverities, t.FilterSources, t.FilterScopes, boolToInt(t.Enabled), boolToInt(t.NotifyOnResolve), time.Now().Unix(), t.ID, ) if err != nil { @@ -193,7 +193,7 @@ func (s *TriggerStoreImpl) ListChannelsForTrigger(ctx context.Context, triggerID func (s *TriggerStoreImpl) ListTriggersForChannel(ctx context.Context, channelID string) ([]*alert.AlertTrigger, error) { rows, err := s.db.QueryContext(ctx, - `SELECT t.id, t.name, t.filter_severities, t.filter_sources, t.filter_scopes, t.filter_tags, + `SELECT t.id, t.name, t.filter_severities, t.filter_sources, t.filter_scopes, t.enabled, t.notify_on_resolve, t.created_at, t.updated_at FROM alert_triggers t JOIN alert_trigger_channels atc ON atc.trigger_id = t.id @@ -235,7 +235,7 @@ func scanTrigger(row *sql.Row) (*alert.AlertTrigger, error) { var enabled, notifyOnResolve int var createdAt, updatedAt int64 err := row.Scan(&t.ID, &t.Name, &t.FilterSeverities, &t.FilterSources, - &t.FilterScopes, &t.FilterTags, &enabled, ¬ifyOnResolve, &createdAt, &updatedAt) + &t.FilterScopes, &enabled, ¬ifyOnResolve, &createdAt, &updatedAt) if err != nil { return nil, err } @@ -251,7 +251,7 @@ func scanTriggerRow(rows *sql.Rows) (*alert.AlertTrigger, error) { var enabled, notifyOnResolve int var createdAt, updatedAt int64 err := rows.Scan(&t.ID, &t.Name, &t.FilterSeverities, &t.FilterSources, - &t.FilterScopes, &t.FilterTags, &enabled, ¬ifyOnResolve, &createdAt, &updatedAt) + &t.FilterScopes, &enabled, ¬ifyOnResolve, &createdAt, &updatedAt) if err != nil { return nil, err } diff --git a/internal/store/triggers_test.go b/internal/store/triggers_test.go index 55e94525..3568c549 100644 --- a/internal/store/triggers_test.go +++ b/internal/store/triggers_test.go @@ -43,7 +43,6 @@ func setupTriggerTestDB(t *testing.T) (*TriggerStoreImpl, *sql.DB) { filter_severities TEXT NOT NULL DEFAULT '', filter_sources TEXT NOT NULL DEFAULT '', filter_scopes TEXT NOT NULL DEFAULT '', - filter_tags TEXT NOT NULL DEFAULT '', enabled INTEGER NOT NULL DEFAULT 1 CHECK (enabled IN (0,1)), notify_on_resolve INTEGER NOT NULL DEFAULT 1 CHECK (notify_on_resolve IN (0,1)), created_at BIGINT NOT NULL DEFAULT 0, diff --git a/internal/store/uuid_schema.sql b/internal/store/uuid_schema.sql index 86f69fc5..61521aea 100644 --- a/internal/store/uuid_schema.sql +++ b/internal/store/uuid_schema.sql @@ -73,7 +73,6 @@ CREATE TABLE containers ( is_ignored INTEGER NOT NULL DEFAULT 0, alert_severity TEXT NOT NULL DEFAULT 'warning' CHECK(alert_severity IN ('critical','warning','info')), restart_threshold INTEGER NOT NULL DEFAULT 3, - alert_channels TEXT, archived INTEGER NOT NULL DEFAULT 0, first_seen_at BIGINT NOT NULL, last_state_change_at BIGINT NOT NULL, @@ -423,7 +422,6 @@ CREATE TABLE alert_triggers ( filter_severities TEXT NOT NULL DEFAULT '', filter_sources TEXT NOT NULL DEFAULT '', filter_scopes TEXT NOT NULL DEFAULT '', - filter_tags TEXT NOT NULL DEFAULT '', enabled INTEGER NOT NULL DEFAULT 1 CHECK(enabled IN (0,1)), notify_on_resolve INTEGER NOT NULL DEFAULT 1 CHECK(notify_on_resolve IN (0,1)), created_at BIGINT NOT NULL DEFAULT 0, @@ -446,7 +444,6 @@ CREATE TABLE escalation_policies ( active_before_downgrade INTEGER NOT NULL DEFAULT 0, severities_json TEXT NOT NULL DEFAULT '[]', scopes_json TEXT NOT NULL DEFAULT '[]', - tags_json TEXT NOT NULL DEFAULT '[]', levels_json TEXT NOT NULL, created_at BIGINT NOT NULL DEFAULT 0, created_by TEXT, diff --git a/internal/swarm/labels.go b/internal/swarm/labels.go index 9eb1780e..4cbf373c 100644 --- a/internal/swarm/labels.go +++ b/internal/swarm/labels.go @@ -19,7 +19,6 @@ const ( labelMaintIgnore = "maintenant.ignore" labelMaintSeverity = "maintenant.alert.severity" labelMaintThreshold = "maintenant.alert.restart_threshold" - labelMaintChannels = "maintenant.alert.channels" ) // IsSwarmManaged returns true if the container has a Swarm service ID label. @@ -53,9 +52,6 @@ func ApplyServiceLabels(c *cmodel.Container, serviceLabels map[string]string) { c.RestartThreshold = n } } - if v, ok := serviceLabels[labelMaintChannels]; ok && v != "" { - c.AlertChannels = v - } // Stack grouping via com.docker.stack.namespace if stack := serviceLabels[labelStackNamespace]; stack != "" { From 378cf84fccc05411999a47942d8dc3cefffa8f1c Mon Sep 17 00:00:00 2001 From: Benjamin Date: Wed, 30 Sep 2026 19:10:30 +0200 Subject: [PATCH 08/54] fix(agent): keep replays out of live state, re-enroll refused agents, trust MAINTENANT_CA_CERT everywhere Replayed spool events now feed history only: resource samples skip the threshold pipeline, container state and health changes write the timeline from the recorded state at their time without touching the current row or emitting, and a replayed container absent from the inventory is kept archived. An agent whose identity the server revokes or no longer knows now enrolls a fresh identity once with the configured token, or exits non-zero with a message naming MAINTENANT_ENROLLMENT_TOKEN. The stream handler also ends when the receive side closes, so a refusal is seen without waiting for a send. The agent probes each labelled endpoint at its own interval and timeout. MAINTENANT_CA_CERT now applies to the agent gRPC client, webhooks and channels, SMTP STARTTLS, license, OSV, changelog, EOL and registry clients. Half a gRPC TLS keypair stops startup, Community logs why no agent listener starts, and the commercial set wires the multi-host extension again. --- .../components/agents/AgentsPanel.vue | 8 +- .../agents/__tests__/AgentsPanel.spec.ts | 79 +++++ internal/agent/agent.go | 88 +++-- internal/agent/client.go | 69 ++-- internal/agent/identity.go | 24 +- internal/agent/prober.go | 89 +++-- internal/agent/prober_test.go | 183 ++++++++-- internal/agent/refusal_test.go | 315 ++++++++++++++++++ internal/app/app.go | 30 +- internal/app/config.go | 14 + internal/app/config_test.go | 17 + internal/app/flags.go | 2 +- internal/app/grpc_listener_test.go | 94 ++++++ internal/commercial/channels/smtp.go | 6 +- internal/commercial/channels/smtp_test.go | 97 ++++++ internal/commercial/extensions.go | 2 + internal/commercial/license/manager.go | 3 +- internal/commercial/license/trust_test.go | 39 +++ internal/commercial/updates/changelog.go | 3 +- internal/commercial/updates/cve.go | 3 +- internal/commercial/updates/trust_test.go | 51 +++ internal/container/agent_event.go | 57 +++- internal/container/agent_event_test.go | 137 +++++++- internal/container/service.go | 115 +++++-- internal/container/service_test.go | 40 ++- internal/eol/fetch.go | 4 +- internal/eol/fetch_test.go | 24 +- internal/resource/agent_event_test.go | 30 ++ internal/resource/service.go | 6 +- internal/ssrf/ssrf.go | 4 +- internal/ssrf/ssrf_test.go | 23 ++ internal/trust/trust.go | 19 +- internal/trust/trust_test.go | 22 ++ internal/trust/trusttest/trusttest.go | 28 ++ internal/update/registry.go | 20 +- internal/update/registry_test.go | 28 ++ 36 files changed, 1557 insertions(+), 216 deletions(-) create mode 100644 frontend/src/commercial/components/agents/__tests__/AgentsPanel.spec.ts create mode 100644 internal/agent/refusal_test.go create mode 100644 internal/app/grpc_listener_test.go create mode 100644 internal/commercial/license/trust_test.go create mode 100644 internal/commercial/updates/trust_test.go create mode 100644 internal/trust/trusttest/trusttest.go create mode 100644 internal/update/registry_test.go diff --git a/frontend/src/commercial/components/agents/AgentsPanel.vue b/frontend/src/commercial/components/agents/AgentsPanel.vue index 376eb88c..57d22d5f 100644 --- a/frontend/src/commercial/components/agents/AgentsPanel.vue +++ b/frontend/src/commercial/components/agents/AgentsPanel.vue @@ -204,14 +204,14 @@ function runtimeLabel(rt: string): string { v-if="agent.spool?.draining" class="ml-1 rounded-full px-2 py-0.5 text-xs font-medium" :style="{ backgroundColor: 'var(--mnt-status-warn-bg)', color: 'var(--mnt-status-warn-text)' }" - :title="`${agent.spool.queued} événements en attente de rejeu`" - >rattrapage · {{ agent.spool.queued }} + :title="`${agent.spool.queued} events waiting to be replayed`" + >catching up · {{ agent.spool.queued }} {{ agent.spool.dropped_since_connect }} perdus + title="Events dropped because the agent's spool was full" + >{{ agent.spool.dropped_since_connect }} lost {{ formatDate(agent.last_seen_at) }} diff --git a/frontend/src/commercial/components/agents/__tests__/AgentsPanel.spec.ts b/frontend/src/commercial/components/agents/__tests__/AgentsPanel.spec.ts new file mode 100644 index 00000000..c39ccd02 --- /dev/null +++ b/frontend/src/commercial/components/agents/__tests__/AgentsPanel.spec.ts @@ -0,0 +1,79 @@ +// Copyright 2026 Benjamin Touchard (Kolapsis) +// SPDX-License-Identifier: LicenseRef-Maintenant-Commercial +// See internal/commercial/LICENSE. + +import { describe, it, expect, vi } from 'vitest' +import { shallowMount } from '@vue/test-utils' +import { ref } from 'vue' +import AgentsPanel from '@/commercial/components/agents/AgentsPanel.vue' +import type { Agent } from '@/services/agentApi' + +const agent: Agent = { + agent_id: 'a1', + hostname: 'web-1', + label: '', + os_arch: 'linux/amd64', + agent_version: '1.0.0', + detected_runtime: 'docker', + status: 'active', + connection_state: 'connected', + last_seen_at: null, + created_at: '2026-09-30T00:00:00Z', + revoked_at: null, + revoked_by: null, + spool: { queued: 42, draining: true, dropped_since_connect: 7, reported_at: '2026-09-30T00:00:00Z' }, + os: { + id: 'debian', + version_id: '12', + pretty_name: 'Debian 12', + source: 'host_file', + unavailable_reason: '', + reported_at: null, + support: { + state: 'supported', + product: 'debian', + cycle: '12', + active_until: null, + security_until: null, + extended_until: null, + days_remaining: null, + table_source: 'embedded', + }, + }, +} + +vi.mock('@/composables/useEdition', () => ({ + useEdition: () => ({ + getQuota: () => ref({ used: 1, limit: -1, isUnlimited: true, isAtLimit: false, nearLimit: false }), + }), +})) + +vi.mock('@/stores/agents', () => ({ + useAgentsStore: () => ({ + agents: [agent], + tokens: [], + metrics: null, + loading: false, + error: null, + fetchAgents: vi.fn(), + fetchTokens: vi.fn(), + fetchMetrics: vi.fn(), + connectSSE: vi.fn(), + disconnectSSE: vi.fn(), + createToken: vi.fn(), + deleteToken: vi.fn(), + }), +})) + +describe('AgentsPanel spool badges', () => { + it('speak the language of the rest of the panel', () => { + const wrapper = shallowMount(AgentsPanel) + const text = wrapper.text() + + expect(text).toContain('catching up · 42') + expect(text).toContain('7 lost') + expect(wrapper.find('[title="42 events waiting to be replayed"]').exists()).toBe(true) + expect(wrapper.find('[title="Events dropped because the agent\'s spool was full"]').exists()).toBe(true) + expect(text).not.toMatch(/rattrapage|perdus/) + }) +}) diff --git a/internal/agent/agent.go b/internal/agent/agent.go index 558020bf..ef4ff9c7 100644 --- a/internal/agent/agent.go +++ b/internal/agent/agent.go @@ -65,31 +65,72 @@ func Run(ctx context.Context, cfg AgentConfig, logger *slog.Logger) error { // Let the server read logs of containers only this host can see. grpcClient.EnableCommands(rt, cfg.AgentVersion, logger) - if !id.Registered { - if cfg.EnrollmentToken == "" { - return fmt.Errorf("agent is not enrolled and --enrollment-token is empty") + enroll := func(ctx context.Context, id *Identity) error { + return RunEnrollment(ctx, id, cfg.DataDir, cfg.EnrollmentToken, rtLabel, cfg.Label, cfg.AgentVersion, grpcClient) + } + serve := func(ctx context.Context, id *Identity) error { + return runEnrolled(ctx, cfg, id, rt, rtLabel, grpcClient, logger) + } + return enrollAndServe(ctx, cfg, id, enroll, serve, logger) +} + +// enrollAndServe enrolls id when needed and serves it, enrolling a fresh identity once with the configured token when the server refuses the stored one. +func enrollAndServe( + ctx context.Context, + cfg AgentConfig, + id *Identity, + enroll func(context.Context, *Identity) error, + serve func(context.Context, *Identity) error, + logger *slog.Logger, +) error { + reenrolled := false + for { + if !id.Registered { + if cfg.EnrollmentToken == "" { + return fmt.Errorf("agent is not enrolled and --enrollment-token is empty") + } + if err := enroll(ctx, id); err != nil { + if reenrolled { + return fmt.Errorf("enrolling again after the server refused the previous identity: %w; "+ + "the stored identity is kept, create a new enrollment token and restart the agent with it", err) + } + return fmt.Errorf("enrollment: %w", err) + } + logger.Info("agent enrolled successfully", "agent_id", id.AgentID) + } else { + logger.Info("agent already enrolled", "agent_id", id.AgentID) } - if err := RunEnrollment(ctx, id, cfg.DataDir, cfg.EnrollmentToken, rtLabel, cfg.Label, cfg.AgentVersion, grpcClient); err != nil { - return fmt.Errorf("enrollment: %w", err) + + err := serve(ctx, id) + if !errors.Is(err, ErrAgentRevokedServer) && !errors.Is(err, ErrAgentUnknownServer) { + return err + } + if cfg.EnrollmentToken == "" || reenrolled { + return fmt.Errorf("%w (agent %s): create an enrollment token on the server and restart the agent "+ + "with MAINTENANT_ENROLLMENT_TOKEN or --enrollment-token set to it", err, id.AgentID) + } + + logger.Warn("agent: the server refused this identity, enrolling a new one with the configured token", + "agent_id", id.AgentID, "reason", err.Error()) + fresh, gerr := newIdentity() + if gerr != nil { + return gerr } - logger.Info("agent enrolled successfully", "agent_id", id.AgentID, "runtime", rtLabel) - } else { - logger.Info("agent already enrolled", "agent_id", id.AgentID) + id = fresh + reenrolled = true } +} + +// runEnrolled streams as id until ctx ends or the server refuses the identity. +func runEnrolled(parent context.Context, cfg AgentConfig, id *Identity, rt runtime.Runtime, rtLabel string, grpcClient *Client, logger *slog.Logger) error { + ctx, cancel := context.WithCancel(parent) + defer cancel() // From here the agent is enrolled and about to stream: report liveness so the // container healthcheck has something to read (the agent serves no HTTP). - // Waiting on the reporter before returning keeps the data directory from - // being written to after the agent has handed control back. healthStopped := StartHealthReporter(ctx, cfg.DataDir, HealthInterval, logger) - defer func() { <-healthStopped }() spool := NewSpool(cfg.DataDir, spoolConfig(cfg), logger) - defer func() { - if cerr := spool.Close(); cerr != nil { - logger.Warn("agent: spool did not close cleanly", "error", cerr) - } - }() if perr := spool.PurgeExpired(ctx); perr != nil { logger.Warn("agent: cannot purge expired spooled events", "error", perr) } @@ -105,20 +146,27 @@ func Run(ctx context.Context, cfg AgentConfig, logger *slog.Logger) error { }() hooks := StreamHooks{Acked: spool.Acked, RateLimited: spool.RateLimited} - err = RunWithReconnect(ctx, grpcClient, id, logger, hooks, func(ctx context.Context, stream *PushStream) error { + err := RunWithReconnect(ctx, grpcClient, id, logger, hooks, func(ctx context.Context, stream *PushStream) error { logger.Info("agent: stream authenticated, draining spool", "agent_id", id.AgentID) spool.ResetDropped() spool.Attach(stream) defer spool.Detach() return spool.Drain(ctx) }) + + // Waiting on the collector and the reporter keeps the data directory from + // being written to after the agent has handed control back. + cancel() <-collectorDone + <-healthStopped - if errors.Is(err, ErrAgentRevokedServer) { + if errors.Is(err, ErrAgentRevokedServer) || errors.Is(err, ErrAgentUnknownServer) { if derr := spool.Discard(); derr != nil { - logger.Warn("agent: cannot discard spool after revocation", "error", derr) + logger.Warn("agent: cannot discard spool after the server refused the identity", "error", derr) } - return fmt.Errorf("agent has been revoked by the server — re-enroll to reconnect") + } + if cerr := spool.Close(); cerr != nil { + logger.Warn("agent: spool did not close cleanly", "error", cerr) } return err } diff --git a/internal/agent/client.go b/internal/agent/client.go index 1cdf88ac..fbe4e2f4 100644 --- a/internal/agent/client.go +++ b/internal/agent/client.go @@ -5,7 +5,6 @@ package agent import ( "context" - "crypto/tls" "errors" "fmt" "log/slog" @@ -23,12 +22,16 @@ import ( "github.com/kolapsis/maintenant/internal/agentpb" "github.com/kolapsis/maintenant/internal/retry" "github.com/kolapsis/maintenant/internal/runtime" + "github.com/kolapsis/maintenant/internal/trust" ) // ErrAgentRevokedServer is returned by RunWithReconnect when the server revokes the agent. // The caller should exit without retrying. var ErrAgentRevokedServer = errors.New("agent revoked by server") +// ErrAgentUnknownServer is returned by RunWithReconnect when the server has no record of the agent. +var ErrAgentUnknownServer = errors.New("agent unknown to the server, deleted or enrolled with another one") + // streamErrorRateLimited is the code the server sends when an agent pushes // events faster than its per-agent allowance. const streamErrorRateLimited = "rate_limited" @@ -70,10 +73,8 @@ func NewClient(ctx context.Context, serverURL string, insecureSkipVerify bool, l var creds credentials.TransportCredentials if useTLS { - tlsCfg := &tls.Config{ - MinVersion: tls.VersionTLS12, - InsecureSkipVerify: insecureSkipVerify, // #nosec G402 -- explicit opt-in flag, warning logged at boot. - } + tlsCfg := trust.ClientTLSConfig() + tlsCfg.InsecureSkipVerify = insecureSkipVerify // #nosec G402 -- explicit opt-in flag, warning logged at boot. if insecureSkipVerify { logger.Warn("TLS certificate verification is disabled — do not use in production") } @@ -116,6 +117,7 @@ type PushStream struct { mu sync.Mutex stream agentpb.Ingest_PushClient recvCh chan error + done chan struct{} // commands executes server-issued commands; nil disables the command channel. commands *CommandRunner @@ -203,6 +205,7 @@ func (ps *PushStream) recvLoop(logger *slog.Logger) { ps.commands.CancelAll() } ps.recvCh <- retErr + close(ps.done) } // DialPush opens the bidirectional Push stream, performs the Ed25519 auth handshake, @@ -251,6 +254,7 @@ func (c *Client) DialPush(ctx context.Context, id *Identity, logger *slog.Logger ps := &PushStream{ stream: stream, recvCh: make(chan error, 1), + done: make(chan struct{}), commands: c.commands, hooks: hooks, ctx: ctx, @@ -261,7 +265,8 @@ func (c *Client) DialPush(ctx context.Context, id *Identity, logger *slog.Logger } // RunWithReconnect runs onStream with exponential backoff reconnect. -// Returns nil when ctx is cancelled, ErrAgentRevokedServer when the server revokes the agent. +// Returns nil when ctx is cancelled, ErrAgentRevokedServer or ErrAgentUnknownServer +// when the server refuses the identity. // Backoff: min(60s, 1s * 2^attempt) ±25% jitter. Attempt resets to 0 if stream was stable >30s. func RunWithReconnect( ctx context.Context, @@ -285,13 +290,23 @@ func RunWithReconnect( if ctx.Err() != nil { return nil } - if isRevokedErr(dialErr) { - logger.Error("agent: revoked during dial, exiting", "err", dialErr) - return ErrAgentRevokedServer + if refusal := refusedIdentity(dialErr); refusal != nil { + logger.Error("agent: identity refused during dial", "err", dialErr) + return refusal } logger.Warn("agent: push dial failed", "err", dialErr, "attempt", backoff.Attempt()) } else { - streamErr := onStream(ctx, stream) + // A sender only notices a closed stream on its next write, which may never come. + streamCtx, stopStream := context.WithCancel(ctx) + go func() { + select { + case <-stream.done: + stopStream() + case <-streamCtx.Done(): + } + }() + streamErr := onStream(streamCtx, stream) + stopStream() stream.Close() // The server reports revocation on the receive side; the send side only // ever surfaces a generic EOF. Ignoring this error turned a revocation @@ -305,9 +320,9 @@ func RunWithReconnect( if ctx.Err() != nil { return nil } - if isRevokedErr(streamErr) || isRevokedErr(recvErr) { - logger.Error("agent: revoked by server, exiting", "err", errors.Join(streamErr, recvErr)) - return ErrAgentRevokedServer + if refusal := refusedIdentity(errors.Join(streamErr, recvErr)); refusal != nil { + logger.Error("agent: identity refused by server", "err", errors.Join(streamErr, recvErr)) + return refusal } if streamErr != nil || recvErr != nil { logger.Warn("agent: stream closed, will reconnect", @@ -325,13 +340,29 @@ func RunWithReconnect( } } -// isRevokedErr reports whether err is a gRPC PermissionDenied "agent_revoked" status. -func isRevokedErr(err error) bool { - if err == nil { - return false +// refusedIdentity maps a server refusal of the stored identity to its sentinel, or returns nil. +func refusedIdentity(err error) error { + for _, e := range flatten(err) { + var carrier interface{ GRPCStatus() *grpcstatus.Status } + if !errors.As(e, &carrier) { + continue + } + st := carrier.GRPCStatus() + switch { + case st.Code() == codes.PermissionDenied && st.Message() == "agent_revoked": + return ErrAgentRevokedServer + case st.Code() == codes.NotFound && st.Message() == "agent not found": + return ErrAgentUnknownServer + } + } + return nil +} + +func flatten(err error) []error { + if joined, ok := err.(interface{ Unwrap() []error }); ok { + return joined.Unwrap() } - st, ok := grpcstatus.FromError(err) - return ok && st.Code() == codes.PermissionDenied && st.Message() == "agent_revoked" + return []error{err} } // parseServerURL extracts the host:port target and whether TLS should be used. diff --git a/internal/agent/identity.go b/internal/agent/identity.go index 0437bf30..6e1645a4 100644 --- a/internal/agent/identity.go +++ b/internal/agent/identity.go @@ -41,16 +41,9 @@ func LoadOrCreate(dataDir string) (*Identity, error) { return nil, fmt.Errorf("read identity file: %w", err) } - // Generate new keypair - pub, priv, err := ed25519.GenerateKey(rand.Reader) + id, err := newIdentity() if err != nil { - return nil, fmt.Errorf("generate ed25519 key: %w", err) - } - - id := &Identity{ - AgentID: generateAgentID(), - PublicKey: []byte(pub), - PrivateKey: []byte(priv), + return nil, err } encoded, err := json.Marshal(id) // #nosec G117 -- agent identity file persists its own Ed25519 private key by design @@ -76,6 +69,19 @@ func LoadOrCreate(dataDir string) (*Identity, error) { return id, nil } +// newIdentity generates an unregistered identity without persisting it. +func newIdentity() (*Identity, error) { + pub, priv, err := ed25519.GenerateKey(rand.Reader) + if err != nil { + return nil, fmt.Errorf("generate ed25519 key: %w", err) + } + return &Identity{ + AgentID: generateAgentID(), + PublicKey: []byte(pub), + PrivateKey: []byte(priv), + }, nil +} + // Save persists updated identity fields (e.g. Registered) back to disk with mode 0600. func (id *Identity) Save(dataDir string) error { path := filepath.Join(dataDir, identityFile) diff --git a/internal/agent/prober.go b/internal/agent/prober.go index 70cf56cc..d42cdfcc 100644 --- a/internal/agent/prober.go +++ b/internal/agent/prober.go @@ -7,6 +7,7 @@ import ( "context" "fmt" "log/slog" + "reflect" "strconv" "time" @@ -21,9 +22,9 @@ import ( ) const ( - endpointProbeInterval = 30 * time.Second - certScanInterval = 60 * time.Second - certScanTimeout = 10 * time.Second + endpointDiscoveryInterval = 30 * time.Second + certScanInterval = 60 * time.Second + certScanTimeout = 10 * time.Second ) // eventSender pushes an event on the stream; decoupled from *PushStream so the @@ -45,9 +46,8 @@ func runLabelProbers(ctx context.Context, id *Identity, rt runtime.Runtime, spoo g, gCtx := errgroup.WithContext(ctx) g.Go(func() error { - return probeLoop(gCtx, endpointProbeInterval, func() error { - return probeEndpointsOnce(gCtx, id.AgentID, ld, spool.Send, logger) - }) + runEndpointProbes(gCtx, id.AgentID, ld, spool.Send, endpointDiscoveryInterval, logger) + return nil }) g.Go(func() error { return probeLoop(gCtx, certScanInterval, func() error { @@ -78,41 +78,72 @@ func probeLoop(ctx context.Context, interval time.Duration, fn func() error) err } } -// probeEndpointsOnce discovers labelled endpoints across the agent's containers, -// probes each unique target (HTTP/TCP, reusing the server's checkers for parity) -// and pushes an EndpointEvent. The server attaches the result to the monitor it -// provisioned for the same target. -func probeEndpointsOnce(ctx context.Context, agentID string, ld labeledDiscoverer, send eventSender, logger *slog.Logger) error { +// runEndpointProbes probes each labelled endpoint through the server's check +// engine, at the interval and timeout its labels set, until ctx is done. +func runEndpointProbes(ctx context.Context, agentID string, ld labeledDiscoverer, send eventSender, discoverEvery time.Duration, logger *slog.Logger) { + engine := endpoint.NewCheckEngine(func(target string, result endpoint.CheckResult) { + if err := send(endpointEvent(agentID, target, result)); err != nil { + logger.Debug("prober: endpoint result not sent", "target", target, "error", err) + } + }, logger) + defer engine.Stop() + + probing := make(map[string]*endpoint.Endpoint) + refresh := func() { + declared, err := discoverEndpoints(ctx, ld, logger) + if err != nil { + logger.Warn("prober: endpoint discovery failed", "err", err) + return + } + for target, ep := range declared { + if current, ok := probing[target]; ok && current.EndpointType == ep.EndpointType && reflect.DeepEqual(current.Config, ep.Config) { + continue + } + probing[target] = ep + engine.AddEndpoint(ctx, ep) + } + for target := range probing { + if _, ok := declared[target]; !ok { + delete(probing, target) + engine.RemoveEndpoint(target) + } + } + } + + refresh() + t := time.NewTicker(discoverEvery) + defer t.Stop() + for { + select { + case <-ctx.Done(): + return + case <-t.C: + refresh() + } + } +} + +// discoverEndpoints returns the HTTP/TCP endpoints declared on the agent's +// containers, keyed by target; the first declaration of a target wins. +func discoverEndpoints(ctx context.Context, ld labeledDiscoverer, logger *slog.Logger) (map[string]*endpoint.Endpoint, error) { results, err := ld.DiscoverAllWithLabels(ctx) if err != nil { - logger.Warn("prober: endpoint discovery failed", "err", err) - return nil + return nil, err } - seen := make(map[string]bool) + declared := make(map[string]*endpoint.Endpoint) for _, res := range results { parsed, _ := endpoint.ParseEndpointLabels(res.Labels, logger) for _, p := range parsed { - if seen[p.Target] { + if _, seen := declared[p.Target]; seen { continue } - seen[p.Target] = true - - ep := &endpoint.Endpoint{EndpointType: p.EndpointType, Target: p.Target, Config: p.Config} - var result endpoint.CheckResult - switch p.EndpointType { - case endpoint.TypeHTTP: - result = endpoint.CheckHTTP(ctx, ep, logger) - case endpoint.TypeTCP: - result = endpoint.CheckTCP(ctx, ep, logger) - default: + if p.EndpointType != endpoint.TypeHTTP && p.EndpointType != endpoint.TypeTCP { continue } - if err := send(endpointEvent(agentID, p.Target, result)); err != nil { - return fmt.Errorf("send endpoint event: %w", err) - } + declared[p.Target] = &endpoint.Endpoint{ID: p.Target, EndpointType: p.EndpointType, Target: p.Target, Config: p.Config} } } - return nil + return declared, nil } // scanCertsOnce discovers labelled TLS targets, scans each unique host:port diff --git a/internal/agent/prober_test.go b/internal/agent/prober_test.go index a502e06e..5c82e922 100644 --- a/internal/agent/prober_test.go +++ b/internal/agent/prober_test.go @@ -9,6 +9,8 @@ import ( "log/slog" "net/http" "net/http/httptest" + "sync" + "sync/atomic" "testing" "time" @@ -26,13 +28,84 @@ func proberTestLogger() *slog.Logger { // fakeLabeledDiscoverer implements labeledDiscoverer for prober tests. type fakeLabeledDiscoverer struct { + mu sync.Mutex results []*docker.DiscoveryResult } func (f *fakeLabeledDiscoverer) DiscoverAllWithLabels(_ context.Context) ([]*docker.DiscoveryResult, error) { + f.mu.Lock() + defer f.mu.Unlock() return f.results, nil } +func (f *fakeLabeledDiscoverer) set(results []*docker.DiscoveryResult) { + f.mu.Lock() + defer f.mu.Unlock() + f.results = results +} + +type probeLog struct { + mu sync.Mutex + events []*agentpb.AgentEvent +} + +func (l *probeLog) send(ev *agentpb.AgentEvent) error { + l.mu.Lock() + defer l.mu.Unlock() + l.events = append(l.events, ev) + return nil +} + +func (l *probeLog) count(target string) int { + l.mu.Lock() + defer l.mu.Unlock() + n := 0 + for _, ev := range l.events { + if ev.GetEndpoint().GetUrl() == target { + n++ + } + } + return n +} + +func (l *probeLog) all() []*agentpb.AgentEvent { + l.mu.Lock() + defer l.mu.Unlock() + return append([]*agentpb.AgentEvent(nil), l.events...) +} + +// startProbes runs the endpoint prober until the test ends. +func startProbes(t *testing.T, ld labeledDiscoverer, discoverEvery time.Duration) *probeLog { + t.Helper() + log := &probeLog{} + ctx, cancel := context.WithCancel(context.Background()) + done := make(chan struct{}) + go func() { + defer close(done) + runEndpointProbes(ctx, "agent-1", ld, log.send, discoverEvery, proberTestLogger()) + }() + t.Cleanup(func() { + cancel() + <-done + }) + return log +} + +func countingServer(t *testing.T, delay time.Duration) (*httptest.Server, *atomic.Int32) { + t.Helper() + var hits atomic.Int32 + srv := httptest.NewServer(http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) { + hits.Add(1) + select { + case <-time.After(delay): + case <-r.Context().Done(): + } + w.WriteHeader(http.StatusOK) + })) + t.Cleanup(srv.Close) + return srv, &hits +} + func TestEndpointEvent(t *testing.T) { code := 200 up := endpointEvent("agent-1", "http://x", endpoint.CheckResult{Success: true, ResponseTimeMs: 42, HTTPStatus: &code}) @@ -80,57 +153,103 @@ func TestCertEvent_ZeroDatesOmitted(t *testing.T) { assert.Nil(t, ci.GetNotAfter(), "zero NotAfter must not be set") } -func TestProbeEndpointsOnce_HTTP(t *testing.T) { - srv := httptest.NewServer(http.HandlerFunc(func(w http.ResponseWriter, _ *http.Request) { - w.WriteHeader(http.StatusOK) - })) - defer srv.Close() - +func TestRunEndpointProbes_HTTP(t *testing.T) { + srv, _ := countingServer(t, 0) ld := &fakeLabeledDiscoverer{results: []*docker.DiscoveryResult{ {Labels: map[string]string{"maintenant.endpoint.http": srv.URL}}, }} - var got []*agentpb.AgentEvent - send := func(ev *agentpb.AgentEvent) error { got = append(got, ev); return nil } - - err := probeEndpointsOnce(context.Background(), "agent-1", ld, send, proberTestLogger()) - require.NoError(t, err) + log := startProbes(t, ld, time.Hour) - require.Len(t, got, 1) - ep := got[0].GetEndpoint() - require.NotNil(t, ep) - assert.Equal(t, srv.URL, ep.GetUrl()) + require.Eventually(t, func() bool { return log.count(srv.URL) == 1 }, 5*time.Second, 20*time.Millisecond) + got := log.all()[0] + ep := got.GetEndpoint() assert.Equal(t, agentpb.EndpointStatus_ENDPOINT_STATUS_UP, ep.GetStatus()) assert.Equal(t, uint32(200), ep.GetStatusCode()) - assert.Equal(t, "agent-1", got[0].GetAgentId()) + assert.Equal(t, "agent-1", got.GetAgentId()) } -func TestProbeEndpointsOnce_DedupsTargetAcrossContainers(t *testing.T) { - srv := httptest.NewServer(http.HandlerFunc(func(w http.ResponseWriter, _ *http.Request) { - w.WriteHeader(http.StatusOK) - })) - defer srv.Close() - - // Two containers declaring the same endpoint target → probed once. +func TestRunEndpointProbes_DedupsTargetAcrossContainers(t *testing.T) { + srv, hits := countingServer(t, 0) ld := &fakeLabeledDiscoverer{results: []*docker.DiscoveryResult{ {Labels: map[string]string{"maintenant.endpoint.http": srv.URL}}, {Labels: map[string]string{"maintenant.endpoint.http": srv.URL}}, }} - var got []*agentpb.AgentEvent - send := func(ev *agentpb.AgentEvent) error { got = append(got, ev); return nil } + log := startProbes(t, ld, time.Hour) - require.NoError(t, probeEndpointsOnce(context.Background(), "a", ld, send, proberTestLogger())) - assert.Len(t, got, 1, "the same target across containers must be probed once") + require.Eventually(t, func() bool { return log.count(srv.URL) >= 1 }, 5*time.Second, 20*time.Millisecond) + time.Sleep(200 * time.Millisecond) + assert.Equal(t, 1, log.count(srv.URL), "the same target across containers must be probed once") + assert.Equal(t, int32(1), hits.Load()) } -func TestProbeEndpointsOnce_NoLabelsNoEvents(t *testing.T) { +func TestRunEndpointProbes_NoLabelsNoEvents(t *testing.T) { ld := &fakeLabeledDiscoverer{results: []*docker.DiscoveryResult{ {Labels: map[string]string{"unrelated.label": "x"}}, }} - sent := 0 - send := func(*agentpb.AgentEvent) error { sent++; return nil } - require.NoError(t, probeEndpointsOnce(context.Background(), "a", ld, send, proberTestLogger())) - assert.Zero(t, sent, "containers without endpoint labels must not produce events") + log := startProbes(t, ld, 20*time.Millisecond) + + time.Sleep(150 * time.Millisecond) + assert.Empty(t, log.all(), "containers without endpoint labels must not produce events") +} + +func TestRunEndpointProbes_EachEndpointKeepsItsOwnInterval(t *testing.T) { + fast, _ := countingServer(t, 0) + slow, _ := countingServer(t, 0) + ld := &fakeLabeledDiscoverer{results: []*docker.DiscoveryResult{ + {Labels: map[string]string{ + "maintenant.endpoint.0.http": fast.URL, + "maintenant.endpoint.0.interval": "150ms", + "maintenant.endpoint.1.http": slow.URL, + }}, + }} + + log := startProbes(t, ld, time.Hour) + + time.Sleep(1200 * time.Millisecond) + assert.GreaterOrEqual(t, log.count(fast.URL), 4, "an endpoint labelled every 150ms must not wait for a fixed 30s cadence") + assert.Equal(t, 1, log.count(slow.URL), "an endpoint on the default interval is probed once, then every 30s") +} + +func TestRunEndpointProbes_HonoursLabelTimeout(t *testing.T) { + srv, _ := countingServer(t, 3*time.Second) + ld := &fakeLabeledDiscoverer{results: []*docker.DiscoveryResult{ + {Labels: map[string]string{ + "maintenant.endpoint.http": srv.URL, + "maintenant.endpoint.timeout": "300ms", + }}, + }} + + start := time.Now() + log := startProbes(t, ld, time.Hour) + + require.Eventually(t, func() bool { return log.count(srv.URL) == 1 }, 2*time.Second, 20*time.Millisecond, + "a 300ms timeout must end the probe long before the server answers") + assert.Less(t, time.Since(start), 2*time.Second) + assert.Equal(t, agentpb.EndpointStatus_ENDPOINT_STATUS_DOWN, log.all()[0].GetEndpoint().GetStatus()) +} + +func TestRunEndpointProbes_FollowsLabelChanges(t *testing.T) { + kept, _ := countingServer(t, 0) + dropped, _ := countingServer(t, 0) + ld := &fakeLabeledDiscoverer{results: []*docker.DiscoveryResult{ + {Labels: map[string]string{"maintenant.endpoint.http": kept.URL, "maintenant.endpoint.interval": "50ms"}}, + {Labels: map[string]string{"maintenant.endpoint.http": dropped.URL, "maintenant.endpoint.interval": "50ms"}}, + }} + + log := startProbes(t, ld, 50*time.Millisecond) + require.Eventually(t, func() bool { return log.count(dropped.URL) >= 2 }, 5*time.Second, 20*time.Millisecond) + + ld.set([]*docker.DiscoveryResult{ + {Labels: map[string]string{"maintenant.endpoint.http": kept.URL, "maintenant.endpoint.interval": "50ms"}}, + }) + time.Sleep(200 * time.Millisecond) + before := log.count(dropped.URL) + keptBefore := log.count(kept.URL) + time.Sleep(300 * time.Millisecond) + + assert.Equal(t, before, log.count(dropped.URL), "a target whose label is gone must stop being probed") + assert.Greater(t, log.count(kept.URL), keptBefore, "an unchanged target keeps its probe") } diff --git a/internal/agent/refusal_test.go b/internal/agent/refusal_test.go new file mode 100644 index 00000000..eca6f819 --- /dev/null +++ b/internal/agent/refusal_test.go @@ -0,0 +1,315 @@ +// Copyright 2026 Benjamin Touchard (Kolapsis) +// SPDX-License-Identifier: Apache-2.0 + +package agent + +import ( + "context" + "crypto/tls" + "errors" + "fmt" + "io" + "log/slog" + "net" + "net/http/httptest" + "sync" + "sync/atomic" + "testing" + "time" + + "github.com/stretchr/testify/assert" + "github.com/stretchr/testify/require" + "google.golang.org/grpc" + "google.golang.org/grpc/codes" + "google.golang.org/grpc/credentials" + grpcstatus "google.golang.org/grpc/status" + + "github.com/kolapsis/maintenant/internal/agentpb" + "github.com/kolapsis/maintenant/internal/trust/trusttest" +) + +type fakeIngest struct { + agentpb.UnimplementedIngestServer + pushErr error + validToken string + + pushes atomic.Int32 + mu sync.Mutex + enrolled []string +} + +func (f *fakeIngest) RegisterAgent(_ context.Context, req *agentpb.RegisterRequest) (*agentpb.RegisterResponse, error) { + if req.GetEnrollmentToken() != f.validToken { + return nil, grpcstatus.Error(codes.FailedPrecondition, "enrollment token already consumed") + } + f.mu.Lock() + f.enrolled = append(f.enrolled, req.GetAgentId()) + f.mu.Unlock() + return &agentpb.RegisterResponse{}, nil +} + +func (f *fakeIngest) Push(stream grpc.BidiStreamingServer[agentpb.ClientMessage, agentpb.ServerMessage]) error { + f.pushes.Add(1) + if err := stream.Send(&agentpb.ServerMessage{ + Payload: &agentpb.ServerMessage_Challenge{Challenge: &agentpb.AuthChallenge{Nonce: make([]byte, 32)}}, + }); err != nil { + return err + } + if _, err := stream.Recv(); err != nil { + return err + } + return f.pushErr +} + +func (f *fakeIngest) enrolledIDs() []string { + f.mu.Lock() + defer f.mu.Unlock() + return append([]string(nil), f.enrolled...) +} + +func quietLogger() *slog.Logger { + return slog.New(slog.NewTextHandler(io.Discard, nil)) +} + +// startIngest serves srv on loopback, over TLS when cert is set, and returns its agent URL. +func startIngest(t *testing.T, srv agentpb.IngestServer, cert *tls.Certificate) string { + t.Helper() + ln, err := net.Listen("tcp", "127.0.0.1:0") + require.NoError(t, err) + + var opts []grpc.ServerOption + scheme := "grpc://" + if cert != nil { + opts = append(opts, grpc.Creds(credentials.NewTLS(&tls.Config{Certificates: []tls.Certificate{*cert}, MinVersion: tls.VersionTLS12}))) + scheme = "grpcs://" + } + s := grpc.NewServer(opts...) + agentpb.RegisterIngestServer(s, srv) + go func() { _ = s.Serve(ln) }() + t.Cleanup(s.Stop) + return scheme + ln.Addr().String() +} + +func dialIngest(t *testing.T, url string) *Client { + t.Helper() + c, err := NewClient(context.Background(), url, false, quietLogger()) + require.NoError(t, err) + t.Cleanup(func() { _ = c.Close() }) + return c +} + +func TestRefusedIdentity(t *testing.T) { + revoked := grpcstatus.Error(codes.PermissionDenied, "agent_revoked") + unknown := grpcstatus.Error(codes.NotFound, "agent not found") + cases := []struct { + name string + err error + want error + }{ + {"revoked", revoked, ErrAgentRevokedServer}, + {"unknown", unknown, ErrAgentUnknownServer}, + {"wrapped", fmt.Errorf("recv auth challenge: %w", unknown), ErrAgentUnknownServer}, + {"joined after a transport error", errors.Join(io.EOF, revoked), ErrAgentRevokedServer}, + {"session closed", grpcstatus.Error(codes.Unavailable, "session_closed"), nil}, + {"unknown enrollment token", grpcstatus.Error(codes.NotFound, "enrollment token not found"), nil}, + {"demo mode", grpcstatus.Error(codes.PermissionDenied, "demo mode: this server does not accept agents"), nil}, + {"not a status", io.EOF, nil}, + {"nil", nil, nil}, + } + for _, c := range cases { + t.Run(c.name, func(t *testing.T) { + assert.Equal(t, c.want, refusedIdentity(c.err)) + }) + } +} + +// blockUntilDone is a drain with nothing to send: only the receive side can end it. +func blockUntilDone(ctx context.Context, _ *PushStream) error { + <-ctx.Done() + return nil +} + +func TestRunWithReconnect_StopsWhenTheServerRefusesTheIdentity(t *testing.T) { + for _, c := range []struct { + name string + pushErr error + want error + }{ + {"deleted agent", grpcstatus.Error(codes.NotFound, "agent not found"), ErrAgentUnknownServer}, + {"revoked agent", grpcstatus.Error(codes.PermissionDenied, "agent_revoked"), ErrAgentRevokedServer}, + } { + t.Run(c.name, func(t *testing.T) { + srv := &fakeIngest{pushErr: c.pushErr} + client := dialIngest(t, startIngest(t, srv, nil)) + id, err := newIdentity() + require.NoError(t, err) + + ctx, cancel := context.WithTimeout(context.Background(), 10*time.Second) + defer cancel() + err = RunWithReconnect(ctx, client, id, quietLogger(), StreamHooks{}, blockUntilDone) + + require.ErrorIs(t, err, c.want) + assert.Equal(t, int32(1), srv.pushes.Load(), "a refused identity must not be retried") + }) + } +} + +func TestRunWithReconnect_KeepsRetryingAnUnavailableServer(t *testing.T) { + srv := &fakeIngest{pushErr: grpcstatus.Error(codes.Unavailable, "session_closed")} + client := dialIngest(t, startIngest(t, srv, nil)) + id, err := newIdentity() + require.NoError(t, err) + + ctx, cancel := context.WithCancel(context.Background()) + done := make(chan error, 1) + go func() { + done <- RunWithReconnect(ctx, client, id, quietLogger(), StreamHooks{}, blockUntilDone) + }() + + require.Eventually(t, func() bool { return srv.pushes.Load() >= 2 }, 10*time.Second, 50*time.Millisecond, + "an outage must be retried with the backoff") + cancel() + select { + case err := <-done: + assert.NoError(t, err) + case <-time.After(5 * time.Second): + t.Fatal("RunWithReconnect did not return after cancellation") + } +} + +type serveScript struct { + results []error + served []string +} + +func (s *serveScript) serve(_ context.Context, id *Identity) error { + s.served = append(s.served, id.AgentID) + if len(s.served) > len(s.results) { + return nil + } + return s.results[len(s.served)-1] +} + +func registeredIdentity(t *testing.T, dir string) *Identity { + t.Helper() + id, err := LoadOrCreate(dir) + require.NoError(t, err) + id.Registered = true + require.NoError(t, id.Save(dir)) + return id +} + +func TestEnrollAndServe_ReenrollsWithTheConfiguredToken(t *testing.T) { + for _, refusal := range []error{ErrAgentRevokedServer, ErrAgentUnknownServer} { + t.Run(refusal.Error(), func(t *testing.T) { + dir := t.TempDir() + old := registeredIdentity(t, dir) + srv := &fakeIngest{validToken: "fresh-token"} + client := dialIngest(t, startIngest(t, srv, nil)) + cfg := AgentConfig{DataDir: dir, EnrollmentToken: "fresh-token"} + enroll := func(ctx context.Context, id *Identity) error { + return RunEnrollment(ctx, id, dir, cfg.EnrollmentToken, RuntimeDocker, "", "1.0.0", client) + } + script := &serveScript{results: []error{refusal}} + + require.NoError(t, enrollAndServe(context.Background(), cfg, old, enroll, script.serve, quietLogger())) + + require.Len(t, script.served, 2) + assert.Equal(t, old.AgentID, script.served[0]) + assert.NotEqual(t, old.AgentID, script.served[1], "the refused identity must be replaced") + assert.Equal(t, []string{script.served[1]}, srv.enrolledIDs()) + + stored, err := LoadOrCreate(dir) + require.NoError(t, err) + assert.Equal(t, script.served[1], stored.AgentID, "the new identity must survive a restart") + assert.True(t, stored.Registered) + }) + } +} + +func TestEnrollAndServe_StopsWithoutTokenWhenRefused(t *testing.T) { + dir := t.TempDir() + old := registeredIdentity(t, dir) + enrolls := 0 + enroll := func(context.Context, *Identity) error { enrolls++; return nil } + script := &serveScript{results: []error{ErrAgentRevokedServer}} + + err := enrollAndServe(context.Background(), AgentConfig{DataDir: dir}, old, enroll, script.serve, quietLogger()) + + require.ErrorIs(t, err, ErrAgentRevokedServer) + assert.Contains(t, err.Error(), "MAINTENANT_ENROLLMENT_TOKEN", "the message must say how to recover") + assert.Contains(t, err.Error(), old.AgentID) + assert.Zero(t, enrolls) + assert.Len(t, script.served, 1, "no retry loop on a refused identity") +} + +func TestEnrollAndServe_KeepsStoredIdentityWhenReenrollmentFails(t *testing.T) { + dir := t.TempDir() + old := registeredIdentity(t, dir) + srv := &fakeIngest{validToken: "fresh-token"} + client := dialIngest(t, startIngest(t, srv, nil)) + cfg := AgentConfig{DataDir: dir, EnrollmentToken: "consumed-token"} + enroll := func(ctx context.Context, id *Identity) error { + return RunEnrollment(ctx, id, dir, cfg.EnrollmentToken, RuntimeDocker, "", "1.0.0", client) + } + script := &serveScript{results: []error{ErrAgentUnknownServer}} + + err := enrollAndServe(context.Background(), cfg, old, enroll, script.serve, quietLogger()) + + require.Error(t, err) + assert.Contains(t, err.Error(), "already consumed") + assert.Len(t, script.served, 1) + stored, lerr := LoadOrCreate(dir) + require.NoError(t, lerr) + assert.Equal(t, old.AgentID, stored.AgentID, "a failed re-enrollment must not destroy the stored identity") +} + +func TestEnrollAndServe_ReenrollsOnlyOnce(t *testing.T) { + dir := t.TempDir() + old := registeredIdentity(t, dir) + enrolls := 0 + enroll := func(_ context.Context, id *Identity) error { + enrolls++ + id.Registered = true + return nil + } + script := &serveScript{results: []error{ErrAgentRevokedServer, ErrAgentRevokedServer}} + + err := enrollAndServe(context.Background(), AgentConfig{DataDir: dir, EnrollmentToken: "tok"}, old, enroll, script.serve, quietLogger()) + + require.ErrorIs(t, err, ErrAgentRevokedServer) + assert.Equal(t, 1, enrolls) + assert.Len(t, script.served, 2) +} + +func TestEnrollAndServe_OtherOutcomesPassThrough(t *testing.T) { + dir := t.TempDir() + old := registeredIdentity(t, dir) + boom := errors.New("boom") + enroll := func(context.Context, *Identity) error { t.Fatal("no enrollment expected"); return nil } + + for _, res := range []error{nil, boom} { + script := &serveScript{results: []error{res}} + err := enrollAndServe(context.Background(), AgentConfig{DataDir: dir, EnrollmentToken: "tok"}, old, enroll, script.serve, quietLogger()) + assert.Equal(t, res, err) + assert.Len(t, script.served, 1) + } +} + +func TestNewClient_TrustsTheConfiguredCA(t *testing.T) { + tlsSrv := httptest.NewTLSServer(nil) + cert := tlsSrv.TLS.Certificates[0] + x509Cert := tlsSrv.Certificate() + tlsSrv.Close() + + srv := &fakeIngest{validToken: "tok"} + url := startIngest(t, srv, &cert) + req := &agentpb.RegisterRequest{AgentId: "a", EnrollmentToken: "tok"} + + _, err := dialIngest(t, url).Register(context.Background(), req) + require.Error(t, err, "a server signed by an unknown authority must be refused") + + trusttest.Trust(t, x509Cert) + _, err = dialIngest(t, url).Register(context.Background(), req) + require.NoError(t, err, "a server signed by MAINTENANT_CA_CERT must be trusted") +} diff --git a/internal/app/app.go b/internal/app/app.go index df13fc83..29029368 100644 --- a/internal/app/app.go +++ b/internal/app/app.go @@ -39,6 +39,7 @@ import ( "github.com/kolapsis/maintenant/internal/store" "github.com/kolapsis/maintenant/internal/swarm" "github.com/kolapsis/maintenant/internal/telemetry" + "github.com/kolapsis/maintenant/internal/trust" "github.com/kolapsis/maintenant/internal/update" "github.com/kolapsis/maintenant/internal/webhook" gomcp "github.com/modelcontextprotocol/go-sdk/mcp" @@ -427,7 +428,7 @@ func New(cfg Config, logger *slog.Logger, opts ...Option) (*App, error) { var eolFetcher *eol.Fetcher if !cfg.DisableOSEOLRefresh { eolFetcher = &eol.Fetcher{ - Client: &http.Client{Timeout: 15 * time.Second}, + Client: &http.Client{Timeout: 15 * time.Second, Transport: trust.HTTPTransport()}, UserAgent: "maintenant/" + cfg.Version + " (+https://maintenant.dev)", } } @@ -800,6 +801,9 @@ func (a *App) Start(ctx context.Context) error { if err := a.cfg.ValidateHTTP(); err != nil { return err } + if err := a.cfg.ValidateGRPCTLS(); err != nil { + return err + } // Derived so an early return (e.g. a failed bind below) cancels every // background goroutine started with ctx, instead of leaking them until @@ -937,15 +941,21 @@ func (a *App) Start(ctx context.Context) error { a.startRuntimeSupervisor(ctx) // Agent gRPC server — server/embedded modes only, where multi-host is open. - if a.serveAgents != nil && a.multihostPlanAllowed() && a.cfg.Mode != "agent" { - if err := a.serveAgents(ctx, extpoint.GRPCConfig{ - Listen: a.cfg.MultiHost.GRPCListen, - PublicURL: a.cfg.MultiHost.GRPCPublicURL, - TLSCertFile: a.cfg.MultiHost.TLSCertFile, - TLSKeyFile: a.cfg.MultiHost.TLSKeyFile, - Insecure: a.cfg.MultiHost.InsecureGRPC, - }); err != nil { - return fmt.Errorf("start agent gRPC server: %w", err) + if a.serveAgents != nil && a.cfg.Mode != "agent" { + if a.multihostPlanAllowed() { + if err := a.serveAgents(ctx, extpoint.GRPCConfig{ + Listen: a.cfg.MultiHost.GRPCListen, + PublicURL: a.cfg.MultiHost.GRPCPublicURL, + TLSCertFile: a.cfg.MultiHost.TLSCertFile, + TLSKeyFile: a.cfg.MultiHost.TLSKeyFile, + Insecure: a.cfg.MultiHost.InsecureGRPC, + }); err != nil { + return fmt.Errorf("start agent gRPC server: %w", err) + } + } else { + required := extension.MinEdition(extension.CapMultihost) + a.logger.Info("agent gRPC listener not started: agents need the "+string(required)+" edition or above", + "edition", extension.CurrentEdition(), "required_edition", required, "listen", a.cfg.MultiHost.GRPCListen) } } diff --git a/internal/app/config.go b/internal/app/config.go index 22a8557e..b2834eb3 100644 --- a/internal/app/config.go +++ b/internal/app/config.go @@ -286,6 +286,20 @@ func (c Config) ValidateHTTP() error { return nil } +// ErrGRPCTLSPair refuses half a keypair for the agent gRPC listener. +var ErrGRPCTLSPair = errors.New("agent gRPC TLS needs both a certificate and its key") + +// ValidateGRPCTLS refuses a certificate without its key, or the reverse, rather than serving a self-signed certificate in their place. +func (c Config) ValidateGRPCTLS() error { + switch { + case c.MultiHost.TLSCertFile != "" && c.MultiHost.TLSKeyFile == "": + return fmt.Errorf("%w: MAINTENANT_GRPC_TLS_CERT is set but MAINTENANT_GRPC_TLS_KEY is empty", ErrGRPCTLSPair) + case c.MultiHost.TLSKeyFile != "" && c.MultiHost.TLSCertFile == "": + return fmt.Errorf("%w: MAINTENANT_GRPC_TLS_KEY is set but MAINTENANT_GRPC_TLS_CERT is empty", ErrGRPCTLSPair) + } + return nil +} + // ErrDemoRuntime is returned when a demo build is pointed at anything but a remote Docker endpoint. var ErrDemoRuntime = errors.New("demo mode only monitors a remote Docker endpoint: set DOCKER_HOST=tcp://host:port, and leave KUBERNETES_SERVICE_HOST and KUBECONFIG unset") diff --git a/internal/app/config_test.go b/internal/app/config_test.go index 5df75494..5699bb48 100644 --- a/internal/app/config_test.go +++ b/internal/app/config_test.go @@ -89,6 +89,23 @@ func TestConfigValidateHTTP_MCP(t *testing.T) { }) } +func TestConfigValidateGRPCTLS(t *testing.T) { + pair := func(cert, key string) Config { + return Config{MultiHost: MultiHostConfig{TLSCertFile: cert, TLSKeyFile: key}} + } + + assert.NoError(t, pair("", "").ValidateGRPCTLS(), "no keypair falls back to the documented dev certificate") + assert.NoError(t, pair("/tls/cert.pem", "/tls/key.pem").ValidateGRPCTLS()) + + err := pair("/tls/cert.pem", "").ValidateGRPCTLS() + require.ErrorIs(t, err, ErrGRPCTLSPair) + assert.Contains(t, err.Error(), "MAINTENANT_GRPC_TLS_KEY is empty", "the message must name the missing variable") + + err = pair("", "/tls/key.pem").ValidateGRPCTLS() + require.ErrorIs(t, err, ErrGRPCTLSPair) + assert.Contains(t, err.Error(), "MAINTENANT_GRPC_TLS_CERT is empty") +} + func TestConfigFromEnv_MCPAllowUnauthenticated(t *testing.T) { t.Run("absent by default", func(t *testing.T) { assert.False(t, ConfigFromEnv().MCP.AllowUnauthenticated) diff --git a/internal/app/flags.go b/internal/app/flags.go index 6f086c21..f0ce291e 100644 --- a/internal/app/flags.go +++ b/internal/app/flags.go @@ -364,7 +364,7 @@ func init() { { EnvName: "MAINTENANT_ENROLLMENT_TOKEN", FlagName: "enrollment-token", Type: FlagTypeString, Default: "", Sensitive: true, - Description: "Enrollment token (agent mode, first boot)", + Description: "Enrollment token (agent mode: first boot, or a new enrollment when the server refuses the stored identity)", ApplyTo: func(c *Config, v string) error { c.MultiHost.EnrollmentToken = v; return nil }, }, { diff --git a/internal/app/grpc_listener_test.go b/internal/app/grpc_listener_test.go new file mode 100644 index 00000000..d2bffc81 --- /dev/null +++ b/internal/app/grpc_listener_test.go @@ -0,0 +1,94 @@ +// Copyright 2026 Benjamin Touchard (Kolapsis) +// SPDX-License-Identifier: Apache-2.0 + +package app_test + +import ( + "bytes" + "context" + "log/slog" + "net" + "sync" + "testing" + "time" + + "github.com/stretchr/testify/assert" + "github.com/stretchr/testify/require" + + "github.com/kolapsis/maintenant/internal/app" + "github.com/kolapsis/maintenant/internal/commercial" + "github.com/kolapsis/maintenant/internal/extension" +) + +type logSink struct { + mu sync.Mutex + buf bytes.Buffer +} + +func (s *logSink) Write(p []byte) (int, error) { + s.mu.Lock() + defer s.mu.Unlock() + return s.buf.Write(p) +} + +func (s *logSink) String() string { + s.mu.Lock() + defer s.mu.Unlock() + return s.buf.String() +} + +func TestStart_CommunityExplainsWhyNoAgentListener(t *testing.T) { + withEdition(t, extension.Community) + cfg, _ := modeGateCfg(t, "embedded") + sink := &logSink{} + logger := slog.New(slog.NewTextHandler(sink, &slog.HandlerOptions{Level: slog.LevelInfo})) + + a, err := app.New(cfg, logger, app.WithExtensions(commercial.Extensions())) + require.NoError(t, err) + _ = startAndCollect(t, a, 2*time.Second) + + out := sink.String() + assert.Contains(t, out, "agent gRPC listener not started") + assert.Contains(t, out, "required_edition="+string(extension.MinEdition(extension.CapMultihost))) + assert.Contains(t, out, "edition="+string(extension.Community)) + assert.NotContains(t, out, "agent gRPC server listening") +} + +func TestStart_ServesAgentsWhereMultihostIsOpen(t *testing.T) { + withEdition(t, extension.Pro) + cfg, logger := modeGateCfg(t, "server") + + a, err := app.New(cfg, logger, app.WithExtensions(commercial.Extensions())) + require.NoError(t, err) + + ctx, cancel := context.WithCancel(context.Background()) + done := make(chan error, 1) + go func() { done <- a.Start(ctx) }() + t.Cleanup(func() { + cancel() + <-done + }) + + require.Eventually(t, func() bool { + conn, err := net.DialTimeout("tcp", cfg.MultiHost.GRPCListen, 200*time.Millisecond) + if err != nil { + return false + } + _ = conn.Close() + return true + }, 10*time.Second, 100*time.Millisecond, "the commercial build must serve agents where multi-host is open") +} + +func TestStart_RefusesHalfAGRPCKeypair(t *testing.T) { + withEdition(t, extension.Pro) + cfg, logger := modeGateCfg(t, "server") + cfg.MultiHost.InsecureGRPC = false + cfg.MultiHost.TLSCertFile = "/etc/maintenant/grpc.crt" + + a, err := app.New(cfg, logger, app.WithExtensions(commercial.Extensions())) + require.NoError(t, err) + + err = startAndCollect(t, a, 5*time.Second) + require.ErrorIs(t, err, app.ErrGRPCTLSPair, "half a keypair must stop startup instead of serving a self-signed certificate") + assert.Contains(t, err.Error(), "MAINTENANT_GRPC_TLS_KEY") +} diff --git a/internal/commercial/channels/smtp.go b/internal/commercial/channels/smtp.go index e02ef457..eef75e56 100644 --- a/internal/commercial/channels/smtp.go +++ b/internal/commercial/channels/smtp.go @@ -6,12 +6,13 @@ package channels import ( "context" - "crypto/tls" "fmt" "net" "net/smtp" "strings" "time" + + "github.com/kolapsis/maintenant/internal/trust" ) const smtpTimeout = 30 * time.Second @@ -62,7 +63,8 @@ func (s *SMTPSender) Send(ctx context.Context, to, subject, textBody string) err }(c) // STARTTLS best-effort: some servers don't support it, so continue in plaintext on error. - tlsCfg := &tls.Config{ServerName: s.cfg.Host, MinVersion: tls.VersionTLS12} + tlsCfg := trust.ClientTLSConfig() + tlsCfg.ServerName = s.cfg.Host _ = c.StartTLS(tlsCfg) // AUTH PLAIN if credentials are configured diff --git a/internal/commercial/channels/smtp_test.go b/internal/commercial/channels/smtp_test.go index cbe4664e..d4699e2c 100644 --- a/internal/commercial/channels/smtp_test.go +++ b/internal/commercial/channels/smtp_test.go @@ -5,14 +5,21 @@ package channels import ( + "bufio" "context" + "crypto/tls" + "fmt" "net" + "net/http/httptest" + "strings" "sync" "testing" "time" "github.com/stretchr/testify/assert" "github.com/stretchr/testify/require" + + "github.com/kolapsis/maintenant/internal/trust/trusttest" ) // silentServer accepts connections and never says a word. @@ -73,3 +80,93 @@ func TestSMTPSenderDeliversPlainText(t *testing.T) { assert.Contains(t, data.String(), "Content-Type: text/plain; charset=utf-8") assert.Contains(t, data.String(), "https://status.example.com/status/confirm?token=abc") } + +// serveOneMail plays a one-session STARTTLS relay and reports the bodies it accepted over TLS. +func serveOneMail(ln net.Listener, cert tls.Certificate, delivered chan<- string) { + raw, err := ln.Accept() + if err != nil { + return + } + defer func() { _ = raw.Close() }() + + conn := raw + r := bufio.NewReader(conn) + reply := func(line string) { _, _ = fmt.Fprintf(conn, "%s\r\n", line) } + overTLS := false + + reply("220 relay ready") + for { + line, err := r.ReadString('\n') + if err != nil { + return + } + cmd := strings.ToUpper(strings.TrimSpace(line)) + switch { + case strings.HasPrefix(cmd, "EHLO"): + if overTLS { + reply("250 relay") + } else { + reply("250-relay") + reply("250 STARTTLS") + } + case cmd == "STARTTLS": + reply("220 go ahead") + tc := tls.Server(raw, &tls.Config{Certificates: []tls.Certificate{cert}, MinVersion: tls.VersionTLS12}) + if err := tc.Handshake(); err != nil { + return + } + conn, r, overTLS = tc, bufio.NewReader(tc), true + case strings.HasPrefix(cmd, "MAIL FROM"), strings.HasPrefix(cmd, "RCPT TO"): + reply("250 ok") + case cmd == "DATA": + reply("354 end with .") + var body strings.Builder + for { + l, err := r.ReadString('\n') + if err != nil { + return + } + if l == ".\r\n" { + break + } + body.WriteString(l) + } + if overTLS { + delivered <- body.String() + } + reply("250 queued") + case cmd == "QUIT": + reply("221 bye") + return + default: + reply("502 not implemented") + } + } +} + +func TestSMTPSender_StartTLSTrustsTheConfiguredCA(t *testing.T) { + borrowed := httptest.NewTLSServer(nil) + cert := borrowed.TLS.Certificates[0] + ca := borrowed.Certificate() + borrowed.Close() + + ln, err := net.Listen("tcp", "127.0.0.1:0") + require.NoError(t, err) + t.Cleanup(func() { _ = ln.Close() }) + delivered := make(chan string, 1) + go serveOneMail(ln, cert, delivered) + + trusttest.Trust(t, ca) + host, port, err := net.SplitHostPort(ln.Addr().String()) + require.NoError(t, err) + sender := NewSMTPSender(SMTPConfig{Host: host, Port: port, From: "maintenant@example.com"}) + + require.NoError(t, sender.Send(context.Background(), "ops@example.com", "disk full", "the disk is full"), + "a relay signed by MAINTENANT_CA_CERT must be trusted") + select { + case body := <-delivered: + assert.Contains(t, body, "the disk is full") + case <-time.After(5 * time.Second): + t.Fatal("no mail delivered over TLS") + } +} diff --git a/internal/commercial/extensions.go b/internal/commercial/extensions.go index 423e40d3..6d8d3d87 100644 --- a/internal/commercial/extensions.go +++ b/internal/commercial/extensions.go @@ -8,6 +8,7 @@ import ( "github.com/kolapsis/maintenant/internal/commercial/channels" "github.com/kolapsis/maintenant/internal/commercial/escalation" "github.com/kolapsis/maintenant/internal/commercial/maintenance" + "github.com/kolapsis/maintenant/internal/commercial/multihost" "github.com/kolapsis/maintenant/internal/commercial/posture" "github.com/kolapsis/maintenant/internal/commercial/statuspage" "github.com/kolapsis/maintenant/internal/commercial/updates" @@ -23,5 +24,6 @@ func Extensions() extpoint.Set { StatusPage: statuspage.NewStatusPage, Suppressor: maintenance.NewMaintenanceSuppressor, Escalation: escalation.NewEscalation, + MultiHost: multihost.NewMultiHost, } } diff --git a/internal/commercial/license/manager.go b/internal/commercial/license/manager.go index 08b69fd2..e2dafe39 100644 --- a/internal/commercial/license/manager.go +++ b/internal/commercial/license/manager.go @@ -16,6 +16,7 @@ import ( "time" "github.com/kolapsis/maintenant/internal/extension" + "github.com/kolapsis/maintenant/internal/trust" ) const ( @@ -74,7 +75,7 @@ func NewManager(licenseKey, dataDir, version, buildDate string, logger *slog.Log version: version, logger: logger.With("component", "license"), publicKey: pubKey, - client: &http.Client{Timeout: 10 * time.Second}, + client: &http.Client{Timeout: 10 * time.Second, Transport: trust.HTTPTransport()}, stop: make(chan struct{}), } diff --git a/internal/commercial/license/trust_test.go b/internal/commercial/license/trust_test.go new file mode 100644 index 00000000..6e8dbe18 --- /dev/null +++ b/internal/commercial/license/trust_test.go @@ -0,0 +1,39 @@ +// Copyright 2026 Benjamin Touchard (Kolapsis) +// SPDX-License-Identifier: LicenseRef-Maintenant-Commercial +// See internal/commercial/LICENSE. + +package license + +import ( + "crypto/ed25519" + "crypto/rand" + "encoding/base64" + "io" + "log/slog" + "net/http" + "net/http/httptest" + "testing" + + "github.com/stretchr/testify/require" + + "github.com/kolapsis/maintenant/internal/trust/trusttest" +) + +func TestNewManager_TrustsTheConfiguredCA(t *testing.T) { + srv := httptest.NewTLSServer(http.HandlerFunc(func(w http.ResponseWriter, _ *http.Request) {})) + defer srv.Close() + trusttest.Trust(t, srv.Certificate()) + + pub, _, err := ed25519.GenerateKey(rand.Reader) + require.NoError(t, err) + prev := publicKeyB64 + InitPublicKey(base64.StdEncoding.EncodeToString(pub)) + t.Cleanup(func() { publicKeyB64 = prev }) + + m, err := NewManager("key", t.TempDir(), "test", "", slog.New(slog.NewTextHandler(io.Discard, nil))) + require.NoError(t, err) + + resp, err := m.client.Get(srv.URL) + require.NoError(t, err, "a license server behind MAINTENANT_CA_CERT must be trusted") + _ = resp.Body.Close() +} diff --git a/internal/commercial/updates/changelog.go b/internal/commercial/updates/changelog.go index dbce3122..84841bba 100644 --- a/internal/commercial/updates/changelog.go +++ b/internal/commercial/updates/changelog.go @@ -15,6 +15,7 @@ import ( "strings" "time" + "github.com/kolapsis/maintenant/internal/trust" "github.com/kolapsis/maintenant/internal/update" ) @@ -30,7 +31,7 @@ type ChangelogResolver struct { func NewChangelogResolver(registry *update.RegistryClient, logger *slog.Logger) *ChangelogResolver { return &ChangelogResolver{ registry: registry, - client: &http.Client{Timeout: 15 * time.Second}, + client: &http.Client{Timeout: 15 * time.Second, Transport: trust.HTTPTransport()}, logger: logger, token: os.Getenv("GITHUB_TOKEN"), } diff --git a/internal/commercial/updates/cve.go b/internal/commercial/updates/cve.go index 3f3088ca..3a0980ab 100644 --- a/internal/commercial/updates/cve.go +++ b/internal/commercial/updates/cve.go @@ -17,6 +17,7 @@ import ( "strings" "time" + "github.com/kolapsis/maintenant/internal/trust" "github.com/kolapsis/maintenant/internal/update" ) @@ -40,7 +41,7 @@ type CVEClient struct { func NewCVEClient(store update.UpdateStore, logger *slog.Logger) *CVEClient { return &CVEClient{ store: store, - client: &http.Client{Timeout: 30 * time.Second}, + client: &http.Client{Timeout: 30 * time.Second, Transport: trust.HTTPTransport()}, logger: logger, delay: 500 * time.Millisecond, baseURL: osvBaseURL, diff --git a/internal/commercial/updates/trust_test.go b/internal/commercial/updates/trust_test.go new file mode 100644 index 00000000..02bb2037 --- /dev/null +++ b/internal/commercial/updates/trust_test.go @@ -0,0 +1,51 @@ +// Copyright 2026 Benjamin Touchard (Kolapsis) +// SPDX-License-Identifier: LicenseRef-Maintenant-Commercial +// See internal/commercial/LICENSE. + +package updates + +import ( + "context" + "net/http" + "net/http/httptest" + "sync/atomic" + "testing" + + "github.com/stretchr/testify/assert" + "github.com/stretchr/testify/require" + + "github.com/kolapsis/maintenant/internal/trust/trusttest" + "github.com/kolapsis/maintenant/internal/update" +) + +func TestCVEClient_TrustsTheConfiguredCA(t *testing.T) { + var batches atomic.Int32 + mux := http.NewServeMux() + mux.HandleFunc("/v1/querybatch", func(w http.ResponseWriter, _ *http.Request) { + batches.Add(1) + _, _ = w.Write([]byte(`{"results":[{"vulns":[]}]}`)) + }) + srv := httptest.NewTLSServer(mux) + defer srv.Close() + trusttest.Trust(t, srv.Certificate()) + + client := newTestCVEClient(&cveStubStore{}, srv.URL) + _, err := client.QueryCVEs(context.Background(), []ImageCVEQuery{ + {ContainerID: "c1", PackageName: "curl", Ecosystem: "Debian", Version: "7.0"}, + }) + + require.NoError(t, err) + assert.Equal(t, int32(1), batches.Load(), "an OSV mirror behind MAINTENANT_CA_CERT must be reached") +} + +func TestChangelogResolver_TrustsTheConfiguredCA(t *testing.T) { + srv := httptest.NewTLSServer(http.HandlerFunc(func(w http.ResponseWriter, _ *http.Request) {})) + defer srv.Close() + trusttest.Trust(t, srv.Certificate()) + + cr := NewChangelogResolver(update.NewRegistryClient(), testLogger()) + resp, err := cr.client.Get(srv.URL) + + require.NoError(t, err, "a changelog API behind MAINTENANT_CA_CERT must be trusted") + _ = resp.Body.Close() +} diff --git a/internal/container/agent_event.go b/internal/container/agent_event.go index 2d071a25..c3cc90b7 100644 --- a/internal/container/agent_event.go +++ b/internal/container/agent_event.go @@ -67,6 +67,29 @@ func (s *Service) HandleAgentEvent(ctx context.Context, agentID string, ev *agen return s.insertAgentContainer(ctx, agentID, ev, meta) } + if !meta.Replayed { + if err := s.refreshAgentContainer(ctx, c, ev); err != nil { + return err + } + } + + if hs := ev.GetHealthStatus(); hs != "" && (meta.Replayed || c.HealthStatus == nil || string(*c.HealthStatus) != hs) { + h := base + h.Action = "health_status" + h.HealthStatus = hs + s.ProcessEvent(ctx, h) + } + + if action := containerStateToAction(ev.GetState()); action != "" { + st := base + st.Action = action + st.ExitCode = ev.GetStatusMessage() + s.ProcessEvent(ctx, st) + } + return nil +} + +func (s *Service) refreshAgentContainer(ctx context.Context, c *Container, ev *agentpb.ContainerEvent) error { dirty := false if img := ev.GetImage(); img != "" && img != c.Image { c.Image = img @@ -84,24 +107,11 @@ func (s *Service) HandleAgentEvent(ctx context.Context, agentID string, ev *agen c.ArchivedAt = nil dirty = true } - if dirty { - if err := s.store.UpdateContainer(ctx, c); err != nil { - return fmt.Errorf("agent event: update %s: %w", shortID(externalID), err) - } + if !dirty { + return nil } - - if hs := ev.GetHealthStatus(); hs != "" && (c.HealthStatus == nil || string(*c.HealthStatus) != hs) { - h := base - h.Action = "health_status" - h.HealthStatus = hs - s.ProcessEvent(ctx, h) - } - - if action := containerStateToAction(ev.GetState()); action != "" { - st := base - st.Action = action - st.ExitCode = ev.GetStatusMessage() - s.ProcessEvent(ctx, st) + if err := s.store.UpdateContainer(ctx, c); err != nil { + return fmt.Errorf("agent event: update %s: %w", shortID(c.ExternalID), err) } return nil } @@ -200,6 +210,13 @@ func (s *Service) insertAgentContainer(ctx context.Context, agentID string, ev * } c.ID = id + // The fresh inventory precedes any replay, so a replayed container it did not carry is gone. + if meta.Replayed { + if err := s.store.ArchiveContainer(ctx, id, now); err != nil { + return fmt.Errorf("agent event: archive replayed %s: %w", shortID(externalID), err) + } + } + // Record an initial transition so uptime tracking has a starting point, // skipping the no-op created→created case (mirrors Reconcile). if state != StateCreated { @@ -213,6 +230,12 @@ func (s *Service) insertAgentContainer(ctx context.Context, agentID string, ev * } } + if meta.Replayed { + s.logger.Debug("agent event: replayed container no longer reported, kept as history", + "external_id", shortID(externalID), "agent_id", agentID) + return nil + } + s.logger.Info("agent event: container discovered", "external_id", shortID(externalID), "name", c.Name, "agent_id", agentID, "state", string(state)) s.emitEvent(event.ContainerDiscovered, c) diff --git a/internal/container/agent_event_test.go b/internal/container/agent_event_test.go index 82f0e814..5c1f7b97 100644 --- a/internal/container/agent_event_test.go +++ b/internal/container/agent_event_test.go @@ -644,37 +644,150 @@ func TestHandleAgentEvent_InsertErrorIsReturned(t *testing.T) { require.Error(t, err) } -// FR-017: a replayed container event writes the state and the timeline, and -// emits nothing. -func TestHandleAgentEvent_ReplayedRecordsStateWithoutEmitting(t *testing.T) { +// A replayed container event writes the timeline, leaves the current +// state to the fresh inventory, and emits nothing. +func TestHandleAgentEvent_ReplayedStateChangeWritesTimelineOnly(t *testing.T) { store := newSvcStore() var events []capturedEvent - svc := newTestService(store, captureEvents(&events)) + checker := &mockRestartChecker{result: "restart loop"} + svc := newTestService(store, captureEvents(&events), func(d *Deps) { d.RestartChecker = checker }) ctx := context.Background() id := extID("replay") agentID := "agent-replay" seed := makeTestContainer(id, StateRunning) seed.AgentID = agentID + liveChange := time.Now().Add(-time.Minute).Truncate(time.Second) + seed.LastStateChangeAt = liveChange store.seed(seed) - observed := time.Now().Add(-3 * time.Hour).Truncate(time.Second) + died := time.Now().Add(-3 * time.Hour).Truncate(time.Second) + restarted := died.Add(30 * time.Second) + require.NoError(t, svc.HandleAgentEvent(ctx, agentID, &agentpb.ContainerEvent{ + ContainerId: id, Name: "replay", State: agentpb.ContainerState_CONTAINER_STATE_EXITED, StatusMessage: "1", + }, agentevent.Meta{ObservedAt: died, Replayed: true, EventID: "evt-die"})) require.NoError(t, svc.HandleAgentEvent(ctx, agentID, &agentpb.ContainerEvent{ - ContainerId: id, Name: "replay", State: agentpb.ContainerState_CONTAINER_STATE_EXITED, - }, agentevent.Meta{ObservedAt: observed, Replayed: true})) + ContainerId: id, Name: "replay", State: agentpb.ContainerState_CONTAINER_STATE_RUNNING, + }, agentevent.Meta{ObservedAt: restarted, Replayed: true, EventID: "evt-start"})) c, err := store.GetContainerByExternalID(ctx, uid.Agent(agentID), id) require.NoError(t, err) require.NotNil(t, c) - assert.Equal(t, StateExited, c.State, "a replayed event must still move the container state") + assert.Equal(t, StateRunning, c.State, "a replayed event must not overwrite the current state") + assert.True(t, liveChange.Equal(c.LastStateChangeAt), "a replayed event must not move the last live change") transitions := store.transitionsFor(c.ID) - require.Len(t, transitions, 1, "a replayed event must still be written to the timeline") - assert.True(t, observed.Equal(transitions[0].Timestamp)) + require.Len(t, transitions, 2, "both replayed changes belong to the timeline") + assert.Equal(t, StateRunning, transitions[0].PreviousState) + assert.Equal(t, StateExited, transitions[0].NewState) + assert.True(t, died.Equal(transitions[0].Timestamp)) + require.NotNil(t, transitions[0].ExitCode) + assert.Equal(t, 1, *transitions[0].ExitCode) + assert.Equal(t, StateExited, transitions[1].PreviousState) + assert.Equal(t, StateRunning, transitions[1].NewState) + assert.True(t, restarted.Equal(transitions[1].Timestamp)) - for _, e := range events { - assert.NotEqual(t, event.ContainerStateChanged, e.typ, "a replayed event must emit nothing") + assert.Empty(t, events, "a replayed event must neither alert nor notify") + assert.Zero(t, checker.calls, "the restart threshold belongs to live events") +} + +func TestHandleAgentEvent_ReplayedHealthChangeWritesTimelineOnly(t *testing.T) { + store := newSvcStore() + var events []capturedEvent + svc := newTestService(store, captureEvents(&events)) + ctx := context.Background() + id := extID("replayhealth") + + agentID := "agent-replay" + healthy := HealthHealthy + seed := makeTestContainer(id, StateRunning) + seed.AgentID = agentID + seed.HasHealthCheck = true + seed.HealthStatus = &healthy + store.seed(seed) + + sick := time.Now().Add(-2 * time.Hour).Truncate(time.Second) + recovered := sick.Add(time.Minute) + for _, step := range []struct { + health string + at time.Time + id string + }{ + {"unhealthy", sick, "evt-sick"}, + {"healthy", recovered, "evt-recovered"}, + } { + require.NoError(t, svc.HandleAgentEvent(ctx, agentID, &agentpb.ContainerEvent{ + ContainerId: id, Name: "replayhealth", HealthStatus: step.health, + }, agentevent.Meta{ObservedAt: step.at, Replayed: true, EventID: step.id})) } + + stored := store.storedHealthStatus(id) + require.NotNil(t, stored) + assert.Equal(t, HealthHealthy, *stored, "a replayed health change must not overwrite the current health") + + c, err := store.GetContainerByExternalID(ctx, uid.Agent(agentID), id) + require.NoError(t, err) + transitions := store.transitionsFor(c.ID) + require.Len(t, transitions, 2, "a recovery equal to the current health still belongs to the timeline") + require.NotNil(t, transitions[0].PreviousHealth) + assert.Equal(t, HealthHealthy, *transitions[0].PreviousHealth) + assert.Equal(t, HealthUnhealthy, *transitions[0].NewHealth) + assert.True(t, sick.Equal(transitions[0].Timestamp)) + assert.Equal(t, HealthUnhealthy, *transitions[1].PreviousHealth) + assert.Equal(t, HealthHealthy, *transitions[1].NewHealth) + + assert.False(t, hasEvent(events, event.ContainerHealthChanged), "a replayed health change must not reach the health alert") +} + +func TestHandleAgentEvent_ReplayedEventLeavesArchivedContainerArchived(t *testing.T) { + store := newSvcStore() + var events []capturedEvent + svc := newTestService(store, captureEvents(&events)) + ctx := context.Background() + id := extID("gone") + + agentID := "agent-replay" + seed := makeTestContainer(id, StateExited) + seed.AgentID = agentID + seed.Image = "app:2" + seed.Archived = true + store.seed(seed) + + require.NoError(t, svc.HandleAgentEvent(ctx, agentID, &agentpb.ContainerEvent{ + ContainerId: id, Name: "gone", Image: "app:1", State: agentpb.ContainerState_CONTAINER_STATE_RUNNING, + }, agentevent.Meta{ObservedAt: time.Now().Add(-time.Hour), Replayed: true, EventID: "evt-old"})) + + c, err := store.GetContainerByExternalID(ctx, uid.Agent(agentID), id) + require.NoError(t, err) + assert.True(t, c.Archived, "a replayed event must not bring back a container the inventory archived") + assert.Equal(t, "app:2", c.Image, "a replayed event must not roll back the container's metadata") + assert.Equal(t, StateExited, c.State) + assert.Len(t, store.transitionsFor(c.ID), 1, "the replayed start still belongs to the timeline") + assert.Empty(t, events) +} + +func TestHandleAgentEvent_ReplayedUnknownContainerIsKeptAsHistory(t *testing.T) { + store := newSvcStore() + var events []capturedEvent + svc := newTestService(store, captureEvents(&events), func(d *Deps) { + d.AgentRuntime = &mockAgentRuntime{runtime: "docker"} + }) + ctx := context.Background() + id := extID("ephemeral") + + started := time.Now().Add(-time.Hour).Truncate(time.Second) + require.NoError(t, svc.HandleAgentEvent(ctx, "agent-replay", &agentpb.ContainerEvent{ + ContainerId: id, Name: "ephemeral", Image: "job:1", State: agentpb.ContainerState_CONTAINER_STATE_RUNNING, + }, agentevent.Meta{ObservedAt: started, Replayed: true, EventID: "evt-start"})) + + c, err := store.GetContainerByExternalID(ctx, "agent-replay", id) + require.NoError(t, err) + require.NotNil(t, c, "the replayed lifecycle needs a row to hang its timeline on") + assert.True(t, c.Archived, "a container the fresh inventory did not carry is gone") + transitions := store.transitionsFor(c.ID) + require.Len(t, transitions, 1) + assert.True(t, started.Equal(transitions[0].Timestamp)) + assert.False(t, hasEvent(events, event.ContainerDiscovered), "a replayed container must not be announced as discovered") } // At-least-once delivery is by design: a stream that breaks mid-drain resends diff --git a/internal/container/service.go b/internal/container/service.go index 889e9058..e3890e18 100644 --- a/internal/container/service.go +++ b/internal/container/service.go @@ -170,6 +170,26 @@ func (s *Service) lookup(ctx context.Context, evt ContainerEvent) (*Container, e return s.store.GetContainerByExternalID(ctx, uid.Agent(evt.AgentID), evt.ExternalID) } +const replayTimelineDepth = 50 + +// timelineAt returns the state and health the timeline holds for c at ts, falling back to c's current values where it holds none. +func (s *Service) timelineAt(ctx context.Context, c *Container, ts time.Time) (ContainerState, *HealthStatus, error) { + transitions, _, err := s.store.ListTransitionsByContainer(ctx, c.ID, ListTransitionsOpts{Until: &ts, Limit: replayTimelineDepth}) + if err != nil { + return "", nil, err + } + state := c.State + if len(transitions) > 0 { + state = transitions[0].NewState + } + for _, t := range transitions { + if t.NewHealth != nil { + return state, t.NewHealth, nil + } + } + return state, c.HealthStatus, nil +} + func (s *Service) handleStateChange(ctx context.Context, evt ContainerEvent, newState ContainerState) { c, err := s.lookup(ctx, evt) if err != nil { @@ -190,21 +210,30 @@ func (s *Service) handleStateChange(ctx context.Context, evt ContainerEvent, new } previousState := c.State + if evt.Replayed { + previousState, _, err = s.timelineAt(ctx, c, evt.Timestamp) + if err != nil { + s.logger.Error("read timeline for replayed state change", "container_id", c.ID, "error", err) + return + } + } if previousState == newState { s.logger.Debug("container: state unchanged, skipping", "container_id", c.ID, "state", string(previousState)) return } - c.State = newState - c.LastStateChangeAt = evt.Timestamp + if !evt.Replayed { + c.State = newState + c.LastStateChangeAt = evt.Timestamp - if err := s.store.UpdateContainer(ctx, c); err != nil { - s.logger.Error("update container state", "id", c.ID, "error", err) - return - } + if err := s.store.UpdateContainer(ctx, c); err != nil { + s.logger.Error("update container state", "id", c.ID, "error", err) + return + } - s.logger.Info("container: state changed", "container_id", c.ID, "name", c.Name, "previous_state", string(previousState), "new_state", string(newState)) + s.logger.Info("container: state changed", "container_id", c.ID, "name", c.Name, "previous_state", string(previousState), "new_state", string(newState)) + } // Record transition transition := &StateTransition{ @@ -236,6 +265,10 @@ func (s *Service) handleStateChange(ctx context.Context, evt ContainerEvent, new s.logger.Error("insert transition", "container_id", c.ID, "error", err) } + if evt.Replayed { + return + } + // Check restart threshold (T030) // Trigger on any transition back to running from a crash state. // Docker emits die→start (exited→running) during crash-loops; the @@ -244,29 +277,22 @@ func (s *Service) handleStateChange(ctx context.Context, evt ContainerEvent, new result, err := s.restartChecker.Check(ctx, c) if err != nil { s.logger.Error("restart check", "container_id", c.ID, "error", err) - } else if !evt.Replayed { - if result != nil { - s.trackRestartAlert(c.ID) - s.emitEvent(event.ContainerRestartAlert, result) - } else { - // Count is below threshold — emit recovery so the alert engine - // can resolve any previously active restart_loop alert. - s.untrackRestartAlert(c.ID) - s.emitEvent(event.ContainerRestartRecover, map[string]interface{}{ - "container_id": c.ID, - "container_name": c.Name, - "timestamp": evt.Timestamp, - "agent_id": c.AgentID, - }) - } + } else if result != nil { + s.trackRestartAlert(c.ID) + s.emitEvent(event.ContainerRestartAlert, result) + } else { + // Count is below threshold: emit recovery so the alert engine + // can resolve any previously active restart_loop alert. + s.untrackRestartAlert(c.ID) + s.emitEvent(event.ContainerRestartRecover, map[string]interface{}{ + "container_id": c.ID, + "container_name": c.Name, + "timestamp": evt.Timestamp, + "agent_id": c.AgentID, + }) } } - // Replay stays silent, but the state and the transition above are already written. - if evt.Replayed { - return - } - s.emitEvent(event.ContainerStateChanged, map[string]interface{}{ "id": c.ID, "state": newState, @@ -314,22 +340,35 @@ func (s *Service) handleHealthChange(ctx context.Context, evt ContainerEvent) { return } - previousHealth := c.HealthStatus + state, previousHealth := c.State, c.HealthStatus newHealth := HealthStatus(evt.HealthStatus) - s.logger.Debug("container: health changed", "container_id", c.ID, "name", c.Name, "previous_health", previousHealth, "new_health", string(newHealth)) - c.HealthStatus = &newHealth - c.LastStateChangeAt = evt.Timestamp + if evt.Replayed { + state, previousHealth, err = s.timelineAt(ctx, c, evt.Timestamp) + if err != nil { + s.logger.Error("read timeline for replayed health change", "container_id", c.ID, "error", err) + return + } + if previousHealth != nil && *previousHealth == newHealth { + return + } + } + s.logger.Debug("container: health changed", "container_id", c.ID, "name", c.Name, "previous_health", previousHealth, "new_health", string(newHealth), "replayed", evt.Replayed) - if err := s.store.UpdateContainer(ctx, c); err != nil { - s.logger.Error("update container health", "id", c.ID, "error", err) - return + if !evt.Replayed { + c.HealthStatus = &newHealth + c.LastStateChangeAt = evt.Timestamp + + if err := s.store.UpdateContainer(ctx, c); err != nil { + s.logger.Error("update container health", "id", c.ID, "error", err) + return + } } transition := &StateTransition{ ID: evt.recordID("health_transition"), ContainerID: c.ID, - PreviousState: c.State, - NewState: c.State, + PreviousState: state, + NewState: state, PreviousHealth: previousHealth, NewHealth: &newHealth, Timestamp: evt.Timestamp, @@ -338,6 +377,10 @@ func (s *Service) handleHealthChange(ctx context.Context, evt ContainerEvent) { s.logger.Error("insert health transition", "container_id", c.ID, "error", err) } + if evt.Replayed { + return + } + s.emitEvent(event.ContainerHealthChanged, map[string]interface{}{ "id": c.ID, "health_status": newHealth, diff --git a/internal/container/service_test.go b/internal/container/service_test.go index b57645ee..71ba6c26 100644 --- a/internal/container/service_test.go +++ b/internal/container/service_test.go @@ -7,6 +7,7 @@ import ( "context" "errors" "log/slog" + "sort" "sync" "testing" "time" @@ -150,22 +151,47 @@ func (m *svcStore) InsertTransition(_ context.Context, t *StateTransition) (stri return "", m.errInsertTransition } clone := *t - clone.ID = uid.New() + if clone.ID == "" { + clone.ID = uid.New() + } + for i, existing := range m.transitions { + if existing.ID == clone.ID { + m.transitions[i] = &clone + return clone.ID, nil + } + } m.transitions = append(m.transitions, &clone) return clone.ID, nil } -func (m *svcStore) ListTransitionsByContainer(_ context.Context, containerID string, _ ListTransitionsOpts) ([]*StateTransition, int, error) { +// ListTransitionsByContainer mirrors the store: newest first, bounds inclusive, at the second. +func (m *svcStore) ListTransitionsByContainer(_ context.Context, containerID string, opts ListTransitionsOpts) ([]*StateTransition, int, error) { m.mu.Lock() defer m.mu.Unlock() var result []*StateTransition - for _, t := range m.transitions { - if t.ContainerID == containerID { - clone := *t - result = append(result, &clone) + for i := len(m.transitions) - 1; i >= 0; i-- { + t := m.transitions[i] + if t.ContainerID != containerID { + continue } + if opts.Since != nil && t.Timestamp.Unix() < opts.Since.Unix() { + continue + } + if opts.Until != nil && t.Timestamp.Unix() > opts.Until.Unix() { + continue + } + clone := *t + result = append(result, &clone) + } + sort.SliceStable(result, func(i, j int) bool { return result[i].Timestamp.Unix() > result[j].Timestamp.Unix() }) + total := len(result) + if opts.Offset > 0 { + result = result[min(opts.Offset, len(result)):] + } + if opts.Limit > 0 && len(result) > opts.Limit { + result = result[:opts.Limit] } - return result, len(result), nil + return result, total, nil } func (m *svcStore) CountRestartsSince(_ context.Context, _ string, _ time.Time) (int, error) { diff --git a/internal/eol/fetch.go b/internal/eol/fetch.go index 0a23568b..6374a23d 100644 --- a/internal/eol/fetch.go +++ b/internal/eol/fetch.go @@ -10,6 +10,8 @@ import ( "io" "net/http" "time" + + "github.com/kolapsis/maintenant/internal/trust" ) const ( @@ -143,7 +145,7 @@ func (f *Fetcher) client() *http.Client { if f.Client != nil { return f.Client } - return &http.Client{Timeout: fetchTimeout} + return &http.Client{Timeout: fetchTimeout, Transport: trust.HTTPTransport()} } func (f *Fetcher) now() time.Time { diff --git a/internal/eol/fetch_test.go b/internal/eol/fetch_test.go index dfabe68f..e4f39e0e 100644 --- a/internal/eol/fetch_test.go +++ b/internal/eol/fetch_test.go @@ -13,11 +13,20 @@ import ( "strings" "testing" "time" + + "github.com/kolapsis/maintenant/internal/trust/trusttest" ) func fixtureServer(t *testing.T, override map[string]http.HandlerFunc) *httptest.Server { t.Helper() - server := httptest.NewServer(http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) { + server := httptest.NewServer(fixtureHandler(t, override)) + t.Cleanup(server.Close) + return server +} + +func fixtureHandler(t *testing.T, override map[string]http.HandlerFunc) http.HandlerFunc { + t.Helper() + return func(w http.ResponseWriter, r *http.Request) { product := path.Base(r.URL.Path) if handler, ok := override[product]; ok { handler(w, r) @@ -33,9 +42,18 @@ func fixtureServer(t *testing.T, override map[string]http.HandlerFunc) *httptest } w.Header().Set("Content-Type", "application/json") _, _ = w.Write(body) - })) + } +} + +func TestFetch_DefaultClientTrustsTheConfiguredCA(t *testing.T) { + server := httptest.NewTLSServer(fixtureHandler(t, nil)) t.Cleanup(server.Close) - return server + trusttest.Trust(t, server.Certificate()) + + fetcher := &Fetcher{BaseURL: server.URL} + if _, err := fetcher.Fetch(context.Background()); err != nil { + t.Fatalf("a mirror behind MAINTENANT_CA_CERT must be trusted: %v", err) + } } func TestFetch(t *testing.T) { diff --git a/internal/resource/agent_event_test.go b/internal/resource/agent_event_test.go index 75c2f7ad..00ae913e 100644 --- a/internal/resource/agent_event_test.go +++ b/internal/resource/agent_event_test.go @@ -130,3 +130,33 @@ func TestHandleAgentEvent_ReplayedSampleKeepsObservationTime(t *testing.T) { assert.True(t, rstore.snapshots[0].Replayed) assert.False(t, callbackInvoked, "a replayed sample must not feed the threshold pipeline") } + +func TestHandleAgentEvent_ReplayedSamplesNeitherAlertNorMoveBreachCounters(t *testing.T) { + extID := "replay0123456789" + wantID := uid.Container(uid.Agent("agent-r"), extID) + c := &container.Container{ID: wantID, ExternalID: extID, AgentID: "agent-r", Name: "demo"} + + rstore := newMockResourceStore() + live := baseConfig(wantID) + live.CPUConsecutiveBreaches = 1 + rstore.alertConfigs[wantID] = live + + var events []string + svc := newTestService(rstore, buildContainerSvc(newMockContainerStore(c)), func(typ string, _ interface{}) { + events = append(events, typ) + }) + + observed := time.Now().Add(-30 * time.Minute) + for i, eventID := range []string{"evt-1", "evt-2", "evt-3"} { + require.NoError(t, svc.HandleAgentEvent(context.Background(), "agent-r", &agentpb.ResourceSample{ + ContainerId: extID, CpuPercent: 95, MemoryBytes: 95, MemoryLimitBytes: 100, + }, agentevent.Meta{ObservedAt: observed.Add(time.Duration(i) * 10 * time.Second), Replayed: true, EventID: eventID})) + } + + require.Len(t, rstore.snapshots, 3, "replayed samples still feed the history") + assert.Empty(t, events, "a replayed sample must neither alert nor reach the live stream") + cfg := storedConfig(t, rstore, wantID) + assert.Equal(t, AlertStateNormal, cfg.AlertState, "a replayed breach must not open an alert") + assert.Equal(t, 1, cfg.CPUConsecutiveBreaches, "the breach counters belong to live samples") + assert.Equal(t, 0, cfg.MemConsecutiveBreaches) +} diff --git a/internal/resource/service.go b/internal/resource/service.go index d080ba1f..091d98ec 100644 --- a/internal/resource/service.go +++ b/internal/resource/service.go @@ -255,7 +255,11 @@ func (s *Service) processSnapshot(snap *ResourceSnapshot) { return } - if s.eventCallback != nil && !snap.Replayed { + if snap.Replayed { + return + } + + if s.eventCallback != nil { memPercent := 0.0 if snap.MemLimit > 0 { memPercent = float64(snap.MemUsed) / float64(snap.MemLimit) * 100.0 diff --git a/internal/ssrf/ssrf.go b/internal/ssrf/ssrf.go index 500be3a4..bc7c576e 100644 --- a/internal/ssrf/ssrf.go +++ b/internal/ssrf/ssrf.go @@ -18,6 +18,8 @@ import ( "net/url" "syscall" "time" + + "github.com/kolapsis/maintenant/internal/trust" ) // ErrBlockedAddress is returned when a URL resolves to, or a connection targets, @@ -100,7 +102,7 @@ func control(_, address string, _ syscall.RawConn) error { // internal/private IPs on every hop, including redirects. When allowPrivate is // true (dev only, via MAINTENANT_ALLOW_PRIVATE_WEBHOOKS) the guard is disabled. func NewHTTPClient(timeout time.Duration, allowPrivate bool) *http.Client { - transport := http.DefaultTransport.(*http.Transport).Clone() + transport := trust.HTTPTransport() if !allowPrivate { dialer := &net.Dialer{Timeout: 30 * time.Second, KeepAlive: 30 * time.Second, Control: control} transport.DialContext = dialer.DialContext diff --git a/internal/ssrf/ssrf_test.go b/internal/ssrf/ssrf_test.go index 8cf39155..e4eae01b 100644 --- a/internal/ssrf/ssrf_test.go +++ b/internal/ssrf/ssrf_test.go @@ -13,6 +13,9 @@ import ( "github.com/stretchr/testify/assert" "github.com/stretchr/testify/require" + + "github.com/kolapsis/maintenant/internal/trust" + "github.com/kolapsis/maintenant/internal/trust/trusttest" ) func TestIsBlocked(t *testing.T) { @@ -86,3 +89,23 @@ func TestNewHTTPClient_BlocksLoopbackAtDial(t *testing.T) { defer func() { _ = resp.Body.Close() }() assert.Equal(t, http.StatusOK, resp.StatusCode) } + +func TestNewHTTPClient_TrustsTheConfiguredCAWithoutOpeningPrivateRanges(t *testing.T) { + srv := httptest.NewTLSServer(http.HandlerFunc(func(w http.ResponseWriter, _ *http.Request) { + w.WriteHeader(http.StatusOK) + })) + defer srv.Close() + trusttest.Trust(t, srv.Certificate()) + + resp, err := NewHTTPClient(2*time.Second, true).Get(srv.URL) + require.NoError(t, err, "a webhook endpoint signed by MAINTENANT_CA_CERT must be trusted") + _ = resp.Body.Close() + + guarded := NewHTTPClient(2*time.Second, false) + _, err = guarded.Get(srv.URL) + require.Error(t, err) + assert.Contains(t, err.Error(), "ssrf guard", "trusting a CA must not lift the private-range guard") + transport, ok := guarded.Transport.(*http.Transport) + require.True(t, ok) + assert.Same(t, trust.Pool(), transport.TLSClientConfig.RootCAs, "the guarded client verifies against the same roots") +} diff --git a/internal/trust/trust.go b/internal/trust/trust.go index a5a54f82..b1c5c0ce 100644 --- a/internal/trust/trust.go +++ b/internal/trust/trust.go @@ -1,8 +1,9 @@ // Copyright 2026 Benjamin Touchard (Kolapsis) // SPDX-License-Identifier: Apache-2.0 -// Package trust holds the root certificates used to validate the TLS endpoints -// and certificates we monitor. +// Package trust holds the root certificates used to validate every outbound TLS +// connection: monitored endpoints and certificates, the agent's link to its +// server, notification channels and third-party APIs. // // It exists because Go's own escape hatch is a trap for this use case: // SSL_CERT_FILE *replaces* the system bundle rather than extending it, so @@ -13,8 +14,10 @@ package trust import ( + "crypto/tls" "crypto/x509" "fmt" + "net/http" "os" "sync/atomic" ) @@ -59,3 +62,15 @@ func Load(path string) error { func Pool() *x509.CertPool { return pool.Load() } + +// ClientTLSConfig returns a client TLS configuration that verifies servers against Pool. +func ClientTLSConfig() *tls.Config { + return &tls.Config{RootCAs: Pool(), MinVersion: tls.VersionTLS12} +} + +// HTTPTransport returns a clone of http.DefaultTransport that verifies servers against Pool. +func HTTPTransport() *http.Transport { + t := http.DefaultTransport.(*http.Transport).Clone() + t.TLSClientConfig = ClientTLSConfig() + return t +} diff --git a/internal/trust/trust_test.go b/internal/trust/trust_test.go index cad02efe..1589c81c 100644 --- a/internal/trust/trust_test.go +++ b/internal/trust/trust_test.go @@ -7,10 +7,13 @@ import ( "crypto/ecdsa" "crypto/elliptic" "crypto/rand" + "crypto/tls" "crypto/x509" "crypto/x509/pkix" "encoding/pem" "math/big" + "net/http" + "net/http/httptest" "os" "path/filepath" "testing" @@ -111,6 +114,25 @@ func TestLoad_UnreadableFileIsAnError(t *testing.T) { assert.Nil(t, Pool()) } +func TestHTTPTransport_TrustsTheLoadedCA(t *testing.T) { + resetPool(t) + srv := httptest.NewTLSServer(http.HandlerFunc(func(w http.ResponseWriter, _ *http.Request) {})) + defer srv.Close() + + _, err := (&http.Client{Transport: HTTPTransport()}).Get(srv.URL) + require.Error(t, err, "without the bundle the test authority is unknown") + + path := filepath.Join(t.TempDir(), "ca.pem") + require.NoError(t, os.WriteFile(path, pem.EncodeToMemory(&pem.Block{Type: "CERTIFICATE", Bytes: srv.Certificate().Raw}), 0o600)) + require.NoError(t, Load(path)) + + resp, err := (&http.Client{Transport: HTTPTransport()}).Get(srv.URL) + require.NoError(t, err) + _ = resp.Body.Close() + assert.Same(t, Pool(), ClientTLSConfig().RootCAs) + assert.Equal(t, uint16(tls.VersionTLS12), ClientTLSConfig().MinVersion) +} + func TestLoad_NonCertificateContentIsAnError(t *testing.T) { resetPool(t) diff --git a/internal/trust/trusttest/trusttest.go b/internal/trust/trusttest/trusttest.go new file mode 100644 index 00000000..720b3c0b --- /dev/null +++ b/internal/trust/trusttest/trusttest.go @@ -0,0 +1,28 @@ +// Copyright 2026 Benjamin Touchard (Kolapsis) +// SPDX-License-Identifier: Apache-2.0 + +// Package trusttest adds test certificates to the trust pool. +package trusttest + +import ( + "crypto/x509" + "encoding/pem" + "os" + "path/filepath" + "testing" + + "github.com/kolapsis/maintenant/internal/trust" +) + +// Trust loads cert into the trust pool through MAINTENANT_CA_CERT's load path, until the test ends. +func Trust(t testing.TB, cert *x509.Certificate) { + t.Helper() + path := filepath.Join(t.TempDir(), "ca.pem") + if err := os.WriteFile(path, pem.EncodeToMemory(&pem.Block{Type: "CERTIFICATE", Bytes: cert.Raw}), 0o600); err != nil { + t.Fatalf("write CA bundle: %v", err) + } + if err := trust.Load(path); err != nil { + t.Fatalf("load CA bundle: %v", err) + } + t.Cleanup(func() { _ = trust.Load("") }) +} diff --git a/internal/update/registry.go b/internal/update/registry.go index f7df7a73..2cde91b1 100644 --- a/internal/update/registry.go +++ b/internal/update/registry.go @@ -6,6 +6,7 @@ package update import ( "context" "fmt" + "net/http" "runtime" "github.com/google/go-containerregistry/pkg/authn" @@ -13,19 +14,24 @@ import ( v1 "github.com/google/go-containerregistry/pkg/v1" "github.com/google/go-containerregistry/pkg/v1/remote" "github.com/google/go-containerregistry/pkg/v1/types" + + "github.com/kolapsis/maintenant/internal/trust" ) // RegistryClient wraps go-containerregistry for read-only registry operations. -type RegistryClient struct{} +type RegistryClient struct { + transport http.RoundTripper +} // NewRegistryClient creates a new registry client. func NewRegistryClient() *RegistryClient { - return &RegistryClient{} + return &RegistryClient{transport: trust.HTTPTransport()} } -func remoteOptions() []remote.Option { +func (rc *RegistryClient) remoteOptions() []remote.Option { return []remote.Option{ remote.WithAuthFromKeychain(authn.DefaultKeychain), + remote.WithTransport(rc.transport), } } @@ -35,7 +41,7 @@ func (rc *RegistryClient) ListTags(ctx context.Context, imageRef string) ([]stri if err != nil { return nil, fmt.Errorf("parse repository %q: %w", imageRef, err) } - tags, err := remote.List(repo, remoteOptions()...) + tags, err := remote.List(repo, rc.remoteOptions()...) if err != nil { return nil, fmt.Errorf("list tags for %q: %w", imageRef, err) } @@ -49,7 +55,7 @@ func (rc *RegistryClient) GetDigest(ctx context.Context, imageRef string) (strin if err != nil { return "", fmt.Errorf("parse reference %q: %w", imageRef, err) } - desc, err := remote.Get(ref, remoteOptions()...) + desc, err := remote.Get(ref, rc.remoteOptions()...) if err != nil { return "", fmt.Errorf("get manifest for %q: %w", imageRef, err) } @@ -73,7 +79,7 @@ func (rc *RegistryClient) GetManifest(ctx context.Context, imageRef string) (*re if err != nil { return nil, fmt.Errorf("parse reference %q: %w", imageRef, err) } - desc, err := remote.Get(ref, remoteOptions()...) + desc, err := remote.Get(ref, rc.remoteOptions()...) if err != nil { return nil, fmt.Errorf("get manifest for %q: %w", imageRef, err) } @@ -86,7 +92,7 @@ func (rc *RegistryClient) GetConfigLabels(ctx context.Context, imageRef string) if err != nil { return nil, fmt.Errorf("parse reference %q: %w", imageRef, err) } - desc, err := remote.Get(ref, remoteOptions()...) + desc, err := remote.Get(ref, rc.remoteOptions()...) if err != nil { return nil, fmt.Errorf("get manifest for %q: %w", imageRef, err) } diff --git a/internal/update/registry_test.go b/internal/update/registry_test.go new file mode 100644 index 00000000..3a705d19 --- /dev/null +++ b/internal/update/registry_test.go @@ -0,0 +1,28 @@ +// Copyright 2026 Benjamin Touchard (kOlapsis) +// SPDX-License-Identifier: Apache-2.0 + +package update + +import ( + "net/http" + "net/http/httptest" + "testing" + + "github.com/stretchr/testify/require" + + "github.com/kolapsis/maintenant/internal/trust/trusttest" +) + +func TestRegistryClient_TrustsTheConfiguredCA(t *testing.T) { + srv := httptest.NewTLSServer(http.HandlerFunc(func(w http.ResponseWriter, _ *http.Request) {})) + defer srv.Close() + trusttest.Trust(t, srv.Certificate()) + + rc := NewRegistryClient() + req, err := http.NewRequest(http.MethodGet, srv.URL+"/v2/", nil) + require.NoError(t, err) + resp, err := rc.transport.RoundTrip(req) + + require.NoError(t, err, "a private registry behind MAINTENANT_CA_CERT must be trusted") + _ = resp.Body.Close() +} From 0c498e5297ec501cca6b8ee7074f45e7b8dec2b4 Mon Sep 17 00:00:00 2001 From: Benjamin Date: Wed, 30 Sep 2026 19:12:46 +0200 Subject: [PATCH 09/54] fix(kubernetes): live events, loss detection and a Docker fallback for stale kubeconfigs - Each event stream builds and starts its own informers after registering them; before, the factory was started empty at connect time and no Kubernetes event ever reached the server. - The stream ends when the API server misses three probes in a row (15 s apart), so the supervisor goes degraded and reconnects as it does for Docker. Any API answer, a refusal included, counts as reachable. - Container events log a shortened external ID without slicing past the end of short Kubernetes IDs such as "ns/pod", which now reach this code. - Auto-detection from a kubeconfig falls back to Docker when that cluster is unreachable; in-cluster and MAINTENANT_RUNTIME=kubernetes still wait. - Helm chart 1.3.0, appVersion 1.8.0. --- deploy/helm/maintenant/Chart.yaml | 4 +- internal/container/service.go | 14 +- internal/container/service_test.go | 8 ++ internal/kubernetes/events.go | 72 +++++++++- internal/kubernetes/events_test.go | 150 +++++++++++++++++++++ internal/kubernetes/runtime.go | 22 ++- internal/runtime/detect.go | 44 +++--- internal/runtime/detect_kubeconfig_test.go | 105 +++++++++++++++ 8 files changed, 372 insertions(+), 47 deletions(-) create mode 100644 internal/kubernetes/events_test.go create mode 100644 internal/runtime/detect_kubeconfig_test.go diff --git a/deploy/helm/maintenant/Chart.yaml b/deploy/helm/maintenant/Chart.yaml index fa93017a..8c013220 100644 --- a/deploy/helm/maintenant/Chart.yaml +++ b/deploy/helm/maintenant/Chart.yaml @@ -2,8 +2,8 @@ apiVersion: v2 name: maintenant description: Self-discovering infrastructure monitoring for Docker and Kubernetes type: application -version: 1.2.0 -appVersion: "1.2.0" +version: 1.3.0 +appVersion: "1.8.0" home: https://github.com/kolapsis/maintenant sources: - https://github.com/kolapsis/maintenant diff --git a/internal/container/service.go b/internal/container/service.go index e3890e18..986aa3a6 100644 --- a/internal/container/service.go +++ b/internal/container/service.go @@ -193,16 +193,16 @@ func (s *Service) timelineAt(ctx context.Context, c *Container, ts time.Time) (C func (s *Service) handleStateChange(ctx context.Context, evt ContainerEvent, newState ContainerState) { c, err := s.lookup(ctx, evt) if err != nil { - s.logger.Error("get container for state change", "external_id", evt.ExternalID[:12], "error", err) + s.logger.Error("get container for state change", "external_id", shortID(evt.ExternalID), "error", err) return } if c == nil { if newState != StateRunning || s.discoverer == nil || uid.Agent(evt.AgentID) != uid.LocalAgent { - s.logger.Debug("unknown container event, skipping", "external_id", evt.ExternalID[:12], "action", evt.Action) + s.logger.Debug("unknown container event, skipping", "external_id", shortID(evt.ExternalID), "action", evt.Action) return } // New container started after initial reconciliation — discover it. - s.logger.Info("new container detected, running reconciliation", "external_id", evt.ExternalID[:12]) + s.logger.Info("new container detected, running reconciliation", "external_id", shortID(evt.ExternalID)) if err := s.Reconcile(ctx, s.discoverer); err != nil { s.logger.Error("on-demand reconciliation failed", "error", err) } @@ -255,7 +255,7 @@ func (s *Service) handleStateChange(ctx context.Context, evt ContainerEvent, new if evt.Action == "die" && s.logFetcher != nil && c.AgentID == uid.LocalAgent { snippet, err := s.logFetcher.FetchLogSnippet(ctx, evt.ExternalID) if err != nil { - s.logger.Warn("fetch log snippet", "external_id", evt.ExternalID[:12], "error", err) + s.logger.Warn("fetch log snippet", "external_id", shortID(evt.ExternalID), "error", err) } else { transition.LogSnippet = snippet } @@ -307,7 +307,7 @@ func (s *Service) handleStateChange(ctx context.Context, evt ContainerEvent, new func (s *Service) handleDestroy(ctx context.Context, evt ContainerEvent) { c, err := s.lookup(ctx, evt) if err != nil { - s.logger.Error("get container for destroy", "external_id", evt.ExternalID[:12], "error", err) + s.logger.Error("get container for destroy", "external_id", shortID(evt.ExternalID), "error", err) return } if c == nil { @@ -317,7 +317,7 @@ func (s *Service) handleDestroy(ctx context.Context, evt ContainerEvent) { now := evt.Timestamp if err := s.store.ArchiveContainer(ctx, c.ID, now); err != nil { - s.logger.Error("archive container", "external_id", evt.ExternalID[:12], "error", err) + s.logger.Error("archive container", "external_id", shortID(evt.ExternalID), "error", err) return } s.untrackRestartAlert(c.ID) @@ -333,7 +333,7 @@ func (s *Service) handleDestroy(ctx context.Context, evt ContainerEvent) { func (s *Service) handleHealthChange(ctx context.Context, evt ContainerEvent) { c, err := s.lookup(ctx, evt) if err != nil { - s.logger.Error("get container for health change", "external_id", evt.ExternalID[:12], "error", err) + s.logger.Error("get container for health change", "external_id", shortID(evt.ExternalID), "error", err) return } if c == nil { diff --git a/internal/container/service_test.go b/internal/container/service_test.go index 71ba6c26..c0fa0ff6 100644 --- a/internal/container/service_test.go +++ b/internal/container/service_test.go @@ -371,6 +371,14 @@ func TestService_ProcessEvent_StartTransitionsToRunning(t *testing.T) { assert.Equal(t, StateRunning, transitions[0].NewState) } +func TestService_ProcessEvent_ShortKubernetesIDOfAnUnknownContainer(t *testing.T) { + svc := newTestService(newSvcStore()) + + assert.NotPanics(t, func() { + svc.ProcessEvent(context.Background(), makeTestEvent("die", "db/pg-0")) + }) +} + func TestService_HandleStateChange_LogFetcherSkippedForRemote(t *testing.T) { // Remote container (AgentID set): logFetcher must NOT be called — it targets // the server's local runtime and cannot read a remote agent's container logs. diff --git a/internal/kubernetes/events.go b/internal/kubernetes/events.go index 58151b3c..06dc4270 100644 --- a/internal/kubernetes/events.go +++ b/internal/kubernetes/events.go @@ -2,6 +2,7 @@ package kubernetes import ( "context" + "errors" "fmt" "time" @@ -9,18 +10,29 @@ import ( "github.com/kolapsis/maintenant/internal/runtime" appsv1 "k8s.io/api/apps/v1" corev1 "k8s.io/api/core/v1" + k8serrors "k8s.io/apimachinery/pkg/api/errors" + metav1 "k8s.io/apimachinery/pkg/apis/meta/v1" + "k8s.io/client-go/informers" + k8s "k8s.io/client-go/kubernetes" "k8s.io/client-go/tools/cache" ) -// streamEvents uses SharedInformerFactory to watch pods and controllers, -// converting events to runtime.RuntimeEvent. +const informerResync = 30 * time.Second + +// streamEvents turns pod and controller changes into runtime events, from informers of its own, until +// ctx ends, the runtime closes or the API server stops answering; the channel then closes. func (r *Runtime) streamEvents(ctx context.Context) <-chan runtime.RuntimeEvent { out := make(chan runtime.RuntimeEvent, 128) - podInformer := r.factory.Core().V1().Pods().Informer() - depInformer := r.factory.Apps().V1().Deployments().Informer() - ssInformer := r.factory.Apps().V1().StatefulSets().Informer() - dsInformer := r.factory.Apps().V1().DaemonSets().Informer() + r.mu.Lock() + clientset := r.clientset + r.mu.Unlock() + factory := informers.NewSharedInformerFactory(clientset, informerResync) + + podInformer := factory.Core().V1().Pods().Informer() + depInformer := factory.Apps().V1().Deployments().Informer() + ssInformer := factory.Apps().V1().StatefulSets().Informer() + dsInformer := factory.Apps().V1().DaemonSets().Informer() emit := func(evt runtime.RuntimeEvent) { select { @@ -162,14 +174,60 @@ func (r *Runtime) streamEvents(ctx context.Context) <-chan runtime.RuntimeEvent }, }) + stop := make(chan struct{}) + factory.Start(stop) go func() { - <-ctx.Done() + r.waitForLoss(ctx, clientset) + close(stop) + // Shutdown returns once no handler can run, so nothing sends on the closed channel. + factory.Shutdown() close(out) }() return out } +// waitForLoss returns when ctx ends, the runtime closes, or probeMisses probes in a row find the API server unreachable. +func (r *Runtime) waitForLoss(ctx context.Context, clientset k8s.Interface) { + ticker := time.NewTicker(r.probeEvery) + defer ticker.Stop() + misses := 0 + for { + select { + case <-ctx.Done(): + return + case <-r.stopCh: + return + case <-ticker.C: + } + err := probeAPIServer(ctx, clientset, r.probeEvery) + if err == nil { + misses = 0 + continue + } + if ctx.Err() != nil { + return + } + misses++ + r.logger.Warn("kubernetes API server unreachable", "error", err, "misses", misses, "of", r.probeMisses) + if misses >= r.probeMisses { + return + } + } +} + +// probeAPIServer fails only when the API server cannot be reached: any answer, a refusal included, proves it is up. +func probeAPIServer(ctx context.Context, clientset k8s.Interface, timeout time.Duration) error { + pctx, cancel := context.WithTimeout(ctx, timeout) + defer cancel() + _, err := clientset.CoreV1().Namespaces().List(pctx, metav1.ListOptions{Limit: 1}) + var status k8serrors.APIStatus + if err != nil && !errors.As(err, &status) { + return err + } + return nil +} + func podToEvent(action string, pod *corev1.Pod) runtime.RuntimeEvent { state, errorDetail := podState(pod) healthStatus := "" diff --git a/internal/kubernetes/events_test.go b/internal/kubernetes/events_test.go new file mode 100644 index 00000000..9f06648e --- /dev/null +++ b/internal/kubernetes/events_test.go @@ -0,0 +1,150 @@ +// Copyright 2026 Benjamin Touchard (kOlapsis) +// SPDX-License-Identifier: Apache-2.0 + +package kubernetes + +import ( + "context" + "errors" + "fmt" + "io" + "log/slog" + "strings" + "sync/atomic" + "testing" + "time" + + "github.com/stretchr/testify/require" + corev1 "k8s.io/api/core/v1" + k8serrors "k8s.io/apimachinery/pkg/api/errors" + metav1 "k8s.io/apimachinery/pkg/apis/meta/v1" + kruntime "k8s.io/apimachinery/pkg/runtime" + "k8s.io/apimachinery/pkg/runtime/schema" + "k8s.io/client-go/kubernetes/fake" + k8stesting "k8s.io/client-go/testing" + + "github.com/kolapsis/maintenant/internal/runtime" +) + +const testProbeEvery = 20 * time.Millisecond + +func streamRuntime(cs *fake.Clientset) *Runtime { + return &Runtime{ + logger: slog.New(slog.NewTextHandler(io.Discard, nil)), + nsFilter: NewNamespaceFilter("", ""), + clientset: cs, + stopCh: make(chan struct{}), + probeEvery: testProbeEvery, + probeMisses: 2, + } +} + +func newStreamRuntime(t *testing.T, cs *fake.Clientset) *Runtime { + t.Helper() + r := streamRuntime(cs) + t.Cleanup(func() { _ = r.Close() }) + return r +} + +// namespaceListFails makes the probe's namespace list fail with err while fail is set. +func namespaceListFails(cs *fake.Clientset, fail *atomic.Bool, err error) { + cs.PrependReactor("list", "namespaces", func(k8stesting.Action) (bool, kruntime.Object, error) { + if fail.Load() { + return true, nil, err + } + return false, nil, nil + }) +} + +func requireOpenFor(t *testing.T, events <-chan runtime.RuntimeEvent, d time.Duration) { + t.Helper() + deadline := time.After(d) + for { + select { + case _, ok := <-events: + require.True(t, ok, "the event stream closed while the API server answered") + case <-deadline: + return + } + } +} + +func requireClosedWithin(t *testing.T, events <-chan runtime.RuntimeEvent, d time.Duration) { + t.Helper() + deadline := time.After(d) + for { + select { + case _, ok := <-events: + if !ok { + return + } + case <-deadline: + t.Fatalf("the event stream is still open after %s", d) + } + } +} + +func TestStreamEvents_DeliversPodsCreatedWhileStreaming(t *testing.T) { + cs := fake.NewClientset() + r := newStreamRuntime(t, cs) + ctx, cancel := context.WithCancel(context.Background()) + defer cancel() + + events := r.StreamEvents(ctx) + + // The fake API server does not replay what happened between the informer's + // list and its watch, so keep creating pods until one comes through live. + deadline := time.After(5 * time.Second) + for i := 0; ; i++ { + name := fmt.Sprintf("live-%d", i) + _, err := cs.CoreV1().Pods("default").Create(ctx, &corev1.Pod{ObjectMeta: metav1.ObjectMeta{Name: name, Namespace: "default"}}, metav1.CreateOptions{}) + require.NoError(t, err) + select { + case evt := <-events: + require.Equal(t, "start", evt.Action) + require.True(t, strings.HasPrefix(evt.ExternalID, "default/live-"), evt.ExternalID) + cancel() + requireClosedWithin(t, events, time.Second) + return + case <-time.After(100 * time.Millisecond): + case <-deadline: + t.Fatal("no event reached the stream: the informers are not running") + } + } +} + +func TestStreamEvents_EndsWhenTheAPIServerIsLost(t *testing.T) { + cs := fake.NewClientset() + var down atomic.Bool + namespaceListFails(cs, &down, errors.New("dial tcp 10.0.0.1:6443: connect: connection refused")) + r := newStreamRuntime(t, cs) + ctx, cancel := context.WithCancel(context.Background()) + defer cancel() + + events := r.StreamEvents(ctx) + requireOpenFor(t, events, 10*testProbeEvery) + + down.Store(true) + requireClosedWithin(t, events, 2*time.Second) +} + +func TestStreamEvents_ARefusalStillProvesTheAPIServerIsUp(t *testing.T) { + cs := fake.NewClientset() + var refusing atomic.Bool + refusing.Store(true) + namespaceListFails(cs, &refusing, k8serrors.NewForbidden(schema.GroupResource{Resource: "namespaces"}, "", errors.New("RBAC"))) + r := newStreamRuntime(t, cs) + ctx, cancel := context.WithCancel(context.Background()) + defer cancel() + + requireOpenFor(t, r.StreamEvents(ctx), 10*testProbeEvery) +} + +func TestStreamEvents_EndsWhenTheRuntimeCloses(t *testing.T) { + r := streamRuntime(fake.NewClientset()) + events := r.StreamEvents(context.Background()) + + require.NoError(t, r.Close()) + + requireClosedWithin(t, events, time.Second) +} diff --git a/internal/kubernetes/runtime.go b/internal/kubernetes/runtime.go index dfc0c389..7209c025 100644 --- a/internal/kubernetes/runtime.go +++ b/internal/kubernetes/runtime.go @@ -16,7 +16,6 @@ import ( "github.com/kolapsis/maintenant/internal/retry" "github.com/kolapsis/maintenant/internal/runtime" metav1 "k8s.io/apimachinery/pkg/apis/meta/v1" - "k8s.io/client-go/informers" k8s "k8s.io/client-go/kubernetes" "k8s.io/client-go/rest" "k8s.io/client-go/tools/clientcmd" @@ -39,9 +38,11 @@ type Runtime struct { nsFilter *NamespaceFilter clientset k8s.Interface metrics metricsv.Interface - factory informers.SharedInformerFactory stopCh chan struct{} + probeEvery time.Duration + probeMisses int + mu sync.Mutex connected bool metricsAvailable bool @@ -60,10 +61,12 @@ type cpuPrev struct { // NewRuntime creates a Kubernetes runtime. Connection is deferred to Connect(). func NewRuntime(logger *slog.Logger, nsFilter *NamespaceFilter) (*Runtime, error) { return &Runtime{ - logger: logger, - nsFilter: nsFilter, - prevCPU: make(map[string]*cpuPrev), - stopCh: make(chan struct{}), + logger: logger, + nsFilter: nsFilter, + prevCPU: make(map[string]*cpuPrev), + stopCh: make(chan struct{}), + probeEvery: 15 * time.Second, + probeMisses: 3, }, nil } @@ -110,8 +113,6 @@ func (r *Runtime) connect(ctx context.Context, config *rest.Config) error { return fmt.Errorf("kubernetes connectivity check failed: %w", err) } - factory := informers.NewSharedInformerFactory(clientset, 30*time.Second) - // Probe metrics-server availability. metricsOK := false if metricsClient != nil { @@ -127,14 +128,9 @@ func (r *Runtime) connect(ctx context.Context, config *rest.Config) error { r.clientset = clientset r.metrics = metricsClient r.metricsAvailable = metricsOK - r.factory = factory r.connected = true r.mu.Unlock() - // Start informers. - factory.Start(r.stopCh) - factory.WaitForCacheSync(r.stopCh) - r.logger.Info("kubernetes runtime connected") return nil } diff --git a/internal/runtime/detect.go b/internal/runtime/detect.go index 80a1ad36..5c091023 100644 --- a/internal/runtime/detect.go +++ b/internal/runtime/detect.go @@ -71,28 +71,14 @@ func DetectWithOverride(ctx context.Context, logger *slog.Logger, override strin // Try KUBECONFIG for out-of-cluster K8s development. if kubeconfig := os.Getenv("KUBECONFIG"); kubeconfig != "" { - if f, ok := factories["kubernetes"]; ok { - logger.Info("detected Kubernetes via KUBECONFIG", "kubeconfig", kubeconfig, "method", "KUBECONFIG") - rt, err := f(ctx, logger) - if err != nil { - logger.Warn("KUBECONFIG present but Kubernetes runtime failed, falling back to Docker", "error", err) - } else { - logger.Info("runtime initialized", "runtime", rt.Name(), "method", "auto_detect_kubeconfig") - return rt, nil - } + if rt := fromKubeconfig(ctx, logger, kubeconfig, "auto_detect_kubeconfig"); rt != nil { + return rt, nil } } else if home, err := os.UserHomeDir(); err == nil { defaultKubeconfig := home + "/.kube/config" if _, err := os.Stat(defaultKubeconfig); err == nil { - if f, ok := factories["kubernetes"]; ok { - logger.Info("detected default kubeconfig", "path", defaultKubeconfig, "method", "default_kubeconfig") - rt, err := f(ctx, logger) - if err != nil { - logger.Warn("default kubeconfig present but Kubernetes runtime failed, falling back to Docker", "error", err) - } else { - logger.Info("runtime initialized", "runtime", rt.Name(), "method", "auto_detect_default_kubeconfig") - return rt, nil - } + if rt := fromKubeconfig(ctx, logger, defaultKubeconfig, "auto_detect_default_kubeconfig"); rt != nil { + return rt, nil } } } @@ -110,6 +96,28 @@ func DetectWithOverride(ctx context.Context, logger *slog.Logger, override strin return nil, fmt.Errorf("no runtime detected; ensure Docker socket is mounted or set MAINTENANT_RUNTIME; registered: %v", registeredNames()) } +// fromKubeconfig returns the Kubernetes runtime when the cluster of kubeconfig answers, nil to let detection move on to Docker. +func fromKubeconfig(ctx context.Context, logger *slog.Logger, kubeconfig, method string) Runtime { + f, ok := factories["kubernetes"] + if !ok { + return nil + } + logger.Info("detected Kubernetes via kubeconfig", "kubeconfig", kubeconfig, "method", method) + rt, err := f(ctx, logger) + if err != nil { + logger.Warn("kubeconfig present but Kubernetes runtime failed, falling back to Docker", "kubeconfig", kubeconfig, "error", err) + return nil + } + if err := rt.TryConnect(ctx); err != nil { + _ = rt.Close() + logger.Warn("kubeconfig present but its cluster is unreachable, falling back to Docker; set MAINTENANT_RUNTIME=kubernetes to wait for that cluster instead", + "kubeconfig", kubeconfig, "error", err) + return nil + } + logger.Info("runtime initialized", "runtime", rt.Name(), "method", method) + return rt +} + func registeredNames() []string { names := make([]string, 0, len(factories)) for n := range factories { diff --git a/internal/runtime/detect_kubeconfig_test.go b/internal/runtime/detect_kubeconfig_test.go new file mode 100644 index 00000000..1ff8c83b --- /dev/null +++ b/internal/runtime/detect_kubeconfig_test.go @@ -0,0 +1,105 @@ +// Copyright 2026 Benjamin Touchard (kOlapsis) +// SPDX-License-Identifier: Apache-2.0 + +package runtime + +import ( + "context" + "errors" + "log/slog" + "os" + "path/filepath" + "testing" +) + +type unreachableRuntime struct { + fakeRuntime + closed bool +} + +func (u *unreachableRuntime) TryConnect(context.Context) error { + return errors.New("dial tcp 10.0.0.1:6443: connect: connection refused") +} + +func (u *unreachableRuntime) Close() error { + u.closed = true + return nil +} + +func registerUnreachableKubernetes() *unreachableRuntime { + rt := &unreachableRuntime{fakeRuntime: fakeRuntime{name: "kubernetes"}} + Register("kubernetes", func(context.Context, *slog.Logger) (Runtime, error) { return rt, nil }) + return rt +} + +func writeKubeconfig(t *testing.T, path string) { + t.Helper() + if err := os.MkdirAll(filepath.Dir(path), 0o700); err != nil { + t.Fatal(err) + } + if err := os.WriteFile(path, []byte("apiVersion: v1"), 0o600); err != nil { + t.Fatal(err) + } +} + +func TestDetect_UnreachableKubeconfigClusterFallsBackToDocker(t *testing.T) { + cases := map[string]func(t *testing.T){ + "KUBECONFIG": func(t *testing.T) { + path := filepath.Join(t.TempDir(), "config") + writeKubeconfig(t, path) + t.Setenv("KUBECONFIG", path) + }, + "default kubeconfig": func(t *testing.T) { + home := t.TempDir() + writeKubeconfig(t, filepath.Join(home, ".kube", "config")) + t.Setenv("HOME", home) + t.Setenv("KUBECONFIG", "") + }, + } + for name, setup := range cases { + t.Run(name, func(t *testing.T) { + resetFactories() + registerFake("docker") + kube := registerUnreachableKubernetes() + t.Setenv("MAINTENANT_RUNTIME", "") + t.Setenv("KUBERNETES_SERVICE_HOST", "") + setup(t) + + rt, err := Detect(context.Background(), slog.Default()) + if err != nil { + t.Fatalf("unexpected error: %v", err) + } + if rt.Name() != "docker" { + t.Fatalf("expected the docker fallback, got %s", rt.Name()) + } + if !kube.closed { + t.Fatal("the abandoned kubernetes runtime was not closed") + } + }) + } +} + +func TestDetect_UnreachableClusterKeptWhenNotGuessedFromAKubeconfig(t *testing.T) { + cases := map[string]map[string]string{ + "MAINTENANT_RUNTIME forced": {"MAINTENANT_RUNTIME": "kubernetes", "KUBERNETES_SERVICE_HOST": ""}, + "in-cluster": {"MAINTENANT_RUNTIME": "", "KUBERNETES_SERVICE_HOST": "10.0.0.1"}, + } + for name, env := range cases { + t.Run(name, func(t *testing.T) { + resetFactories() + registerFake("docker") + registerUnreachableKubernetes() + for k, v := range env { + t.Setenv(k, v) + } + + rt, err := Detect(context.Background(), slog.Default()) + if err != nil { + t.Fatalf("unexpected error: %v", err) + } + if rt.Name() != "kubernetes" { + t.Fatalf("expected kubernetes, got %s", rt.Name()) + } + }) + } +} From 85d72e3fb6b07925c5bd65edaac74e5e64463874 Mon Sep 17 00:00:00 2001 From: Benjamin Date: Wed, 30 Sep 2026 19:14:22 +0200 Subject: [PATCH 10/54] fix(status): answer subscriptions alike and encode mail subjects Subscribing an address already on file failed with a 500 on its unique constraint. The store now upserts a pending subscription: a new or still unconfirmed address gets a fresh confirmation token and a new 24 h window, a confirmed one is left as is. The confirmation email goes out in the background, so POST /status/subscribe gives the same code, body and delay whether the address is new, pending or confirmed, and whether the mail server answers or not. The confirmation_failed code is gone. buildMIME encodes the Subject header as an RFC 2047 encoded word, so accented incident titles and alert messages reach mail clients intact, for the status page and the email channel alike. --- frontend/src/locales/status-page/en.ts | 1 - frontend/src/locales/status-page/fr.ts | 1 - frontend/src/pages/PublicStatusPage.vue | 1 - internal/api/v1/status_mail_test.go | 46 ++++++++ internal/commercial/channels/smtp.go | 3 +- internal/commercial/channels/smtp_test.go | 46 ++++++++ internal/status/announce_test.go | 41 ++++--- internal/status/handler.go | 10 +- internal/status/handler_test.go | 131 ++++++++++++++++------ internal/status/store.go | 2 +- internal/status/subscriber.go | 28 ++--- internal/store/subscribers.go | 20 ++-- internal/store/subscribers_test.go | 87 ++++++++++++++ 13 files changed, 333 insertions(+), 84 deletions(-) create mode 100644 internal/store/subscribers_test.go diff --git a/frontend/src/locales/status-page/en.ts b/frontend/src/locales/status-page/en.ts index ff0aaca5..9e372efc 100644 --- a/frontend/src/locales/status-page/en.ts +++ b/frontend/src/locales/status-page/en.ts @@ -31,7 +31,6 @@ const en = { subscribeInvalidEmail: 'Enter a valid email address.', subscribeRateLimited: 'Too many attempts. Try again later.', subscribeUnavailable: 'Email updates are not available right now.', - subscribeConfirmationFailed: 'The confirmation email could not be sent. Try again later.', subscribeFailed: 'The subscription failed. Try again later.', updatedAt: 'Updated', poweredBy: 'Powered by Maintenant', diff --git a/frontend/src/locales/status-page/fr.ts b/frontend/src/locales/status-page/fr.ts index 9976b8c1..33e176d2 100644 --- a/frontend/src/locales/status-page/fr.ts +++ b/frontend/src/locales/status-page/fr.ts @@ -31,7 +31,6 @@ const fr = { subscribeInvalidEmail: 'Saisissez une adresse e-mail valide.', subscribeRateLimited: 'Trop de tentatives. Réessayez plus tard.', subscribeUnavailable: 'Les mises à jour par e-mail ne sont pas disponibles pour le moment.', - subscribeConfirmationFailed: 'L\'e-mail de confirmation n\'a pas pu être envoyé. Réessayez plus tard.', subscribeFailed: 'L\'abonnement a échoué. Réessayez plus tard.', updatedAt: 'Mis à jour', poweredBy: 'Propulsé par Maintenant', diff --git a/frontend/src/pages/PublicStatusPage.vue b/frontend/src/pages/PublicStatusPage.vue index e8d84325..c18a70f1 100644 --- a/frontend/src/pages/PublicStatusPage.vue +++ b/frontend/src/pages/PublicStatusPage.vue @@ -171,7 +171,6 @@ const subscribeErrors: Record = { invalid_email: 'subscribeInvalidEmail', rate_limited: 'subscribeRateLimited', subscriptions_unavailable: 'subscribeUnavailable', - confirmation_failed: 'subscribeConfirmationFailed', } async function handleSubscribe() { diff --git a/internal/api/v1/status_mail_test.go b/internal/api/v1/status_mail_test.go index da218e54..59c506ce 100644 --- a/internal/api/v1/status_mail_test.go +++ b/internal/api/v1/status_mail_test.go @@ -183,6 +183,52 @@ func TestStatusPageMailFlow(t *testing.T) { mailer.none(t) } +func TestSubscribeRevealsNothingAboutTheAddress(t *testing.T) { + withEdition(t, extension.Pro) + logger := slog.New(slog.NewTextHandler(io.Discard, nil)) + db := storetest.Open(t, logger) + subscribers := store.NewSubscriberStore(db) + mailer := newRecordingMailer() + svc := status.NewService(status.Deps{ + Components: store.NewStatusComponentStore(db), + Logger: logger, + Subscribers: status.NewSubscriberService(subscribers, mailer, "https://status.example.com", logger), + }) + public := http.NewServeMux() + status.NewHandler(svc, http.NotFoundHandler(), logger, ratelimit.New(5.0/3600.0, 5, nil)).Register(public, nil) + subscribe := func() *httptest.ResponseRecorder { + return serve(t, public, http.MethodPost, "/status/subscribe", `{"email":"visitor@example.com"}`) + } + + fresh := subscribe() + require.Equal(t, http.StatusOK, fresh.Code, fresh.Body.String()) + firstLink := confirmLink.FindString(mailer.next(t).body) + require.NotEmpty(t, firstLink) + + pending := subscribe() + secondLink := confirmLink.FindString(mailer.next(t).body) + require.NotEmpty(t, secondLink) + assert.NotEqual(t, firstLink, secondLink, "a pending address gets a fresh link") + + rec := serve(t, public, http.MethodGet, requestURI(t, firstLink), "") + assert.Equal(t, http.StatusBadRequest, rec.Code, "the replaced link no longer confirms") + rec = serve(t, public, http.MethodGet, requestURI(t, secondLink), "") + require.Equal(t, http.StatusOK, rec.Code, rec.Body.String()) + + confirmed := subscribe() + mailer.none(t) + + for name, got := range map[string]*httptest.ResponseRecorder{"pending": pending, "confirmed": confirmed} { + assert.Equal(t, fresh.Code, got.Code, name) + assert.Equal(t, fresh.Body.String(), got.Body.String(), name) + assert.Equal(t, fresh.Header().Get("Content-Type"), got.Header().Get("Content-Type"), name) + } + stats, err := subscribers.GetSubscriberStats(context.Background()) + require.NoError(t, err) + assert.Equal(t, 1, stats.Total) + assert.Equal(t, 1, stats.Confirmed) +} + func TestStatusSmtpTest(t *testing.T) { withEdition(t, extension.Pro) logger := slog.New(slog.NewTextHandler(io.Discard, nil)) diff --git a/internal/commercial/channels/smtp.go b/internal/commercial/channels/smtp.go index eef75e56..77226cc4 100644 --- a/internal/commercial/channels/smtp.go +++ b/internal/commercial/channels/smtp.go @@ -7,6 +7,7 @@ package channels import ( "context" "fmt" + "mime" "net" "net/smtp" "strings" @@ -102,7 +103,7 @@ func buildMIME(from, to, subject, body string) string { var b strings.Builder b.WriteString("From: " + sanitizeHeader(from) + "\r\n") b.WriteString("To: " + sanitizeHeader(to) + "\r\n") - b.WriteString("Subject: " + sanitizeHeader(subject) + "\r\n") + b.WriteString("Subject: " + mime.QEncoding.Encode("utf-8", sanitizeHeader(subject)) + "\r\n") b.WriteString("MIME-Version: 1.0\r\n") b.WriteString("Content-Type: text/plain; charset=utf-8\r\n") b.WriteString("\r\n") diff --git a/internal/commercial/channels/smtp_test.go b/internal/commercial/channels/smtp_test.go index d4699e2c..6bb22591 100644 --- a/internal/commercial/channels/smtp_test.go +++ b/internal/commercial/channels/smtp_test.go @@ -9,6 +9,7 @@ import ( "context" "crypto/tls" "fmt" + "mime" "net" "net/http/httptest" "strings" @@ -70,6 +71,51 @@ func TestSMTPSenderGivesUpWhenItsContextEnds(t *testing.T) { } } +func subjectHeader(t *testing.T, message string) string { + t.Helper() + for _, line := range strings.Split(message, "\n") { + if value, ok := strings.CutPrefix(line, "Subject: "); ok { + return strings.TrimRight(value, "\r") + } + } + t.Fatalf("no Subject header in %q", message) + return "" +} + +func TestBuildMIMEEncodesANonASCIISubject(t *testing.T) { + const subject = "Résolu : panne de la base de données" + raw := subjectHeader(t, buildMIME("maintenant@example.com", "ops@example.com", subject, "body")) + + assert.True(t, strings.HasPrefix(raw, "=?utf-8?q?"), "RFC 2047 encoded word expected, got %q", raw) + for _, r := range raw { + require.Less(t, r, rune(0x80), "the header carries raw non-ASCII bytes: %q", raw) + } + decoded, err := new(mime.WordDecoder).DecodeHeader(raw) + require.NoError(t, err) + assert.Equal(t, subject, decoded) +} + +func TestBuildMIMELeavesAnASCIISubjectReadable(t *testing.T) { + assert.Equal(t, "[major] Database down", + subjectHeader(t, buildMIME("maintenant@example.com", "ops@example.com", "[major] Database down", "body"))) +} + +func TestBuildMIMEKeepsHeaderInjectionOut(t *testing.T) { + msg := buildMIME("maintenant@example.com", "ops@example.com", "Hello\r\nBcc: victim@example.com", "body") + assert.NotContains(t, msg, "\r\nBcc:") +} + +func TestSMTPSenderDeliversAnAccentedSubject(t *testing.T) { + host, port, data, _ := smtpStub(t) + s := NewSMTPSender(SMTPConfig{Host: host, Port: port, From: "maintenant@example.com"}) + + require.NoError(t, s.Send(context.Background(), "visitor@example.com", "Mise à jour : réseau dégradé", "body")) + + decoded, err := new(mime.WordDecoder).DecodeHeader(subjectHeader(t, data.String())) + require.NoError(t, err) + assert.Equal(t, "Mise à jour : réseau dégradé", decoded) +} + func TestSMTPSenderDeliversPlainText(t *testing.T) { host, port, data, rcpt := smtpStub(t) s := NewSMTPSender(SMTPConfig{Host: host, Port: port, From: "maintenant@example.com"}) diff --git a/internal/status/announce_test.go b/internal/status/announce_test.go index c71a9f1c..b8a83dff 100644 --- a/internal/status/announce_test.go +++ b/internal/status/announce_test.go @@ -89,7 +89,7 @@ func newAnnouncingService(t *testing.T, mailer Mailer) (*Service, *recordingNoti func TestAnnounceIncidentEmailsSubscribersOnce(t *testing.T) { pinEdition(t, extension.Pro) - svc, n, b := newAnnouncingService(t, &fakeMailer{}) + svc, n, b := newAnnouncingService(t, newFakeMailer()) inc := &Incident{ ID: "inc-1", Title: "Database down", Severity: SeverityMajor, Status: IncidentInvestigating, @@ -114,7 +114,7 @@ func TestAnnounceIncidentEmailsSubscribersOnce(t *testing.T) { func TestAnnounceIncidentUpdateResolvingSendsOnlyTheResolution(t *testing.T) { pinEdition(t, extension.Pro) - svc, n, b := newAnnouncingService(t, &fakeMailer{}) + svc, n, b := newAnnouncingService(t, newFakeMailer()) inc := &Incident{ID: "inc-1", Title: "Database down", Status: IncidentInvestigating} svc.AnnounceIncidentUpdate(context.Background(), inc, &IncidentUpdate{IncidentID: "inc-1", Status: IncidentResolved, Message: "Back to normal."}) @@ -134,7 +134,7 @@ func TestAnnounceIncidentUpdateResolvingSendsOnlyTheResolution(t *testing.T) { func TestAnnounceIncidentUpdateOnAnOpenIncident(t *testing.T) { pinEdition(t, extension.Pro) - svc, n, b := newAnnouncingService(t, &fakeMailer{}) + svc, n, b := newAnnouncingService(t, newFakeMailer()) inc := &Incident{ID: "inc-1", Title: "Database down", Status: IncidentInvestigating} svc.AnnounceIncidentUpdate(context.Background(), inc, &IncidentUpdate{IncidentID: "inc-1", Status: "monitoring", Message: "Fix deployed."}) @@ -154,7 +154,7 @@ func TestAnnounceIncidentUpdateOnAnOpenIncident(t *testing.T) { func TestAnnounceIncidentUpdateAfterResolutionIsAnUpdate(t *testing.T) { pinEdition(t, extension.Pro) - svc, n, _ := newAnnouncingService(t, &fakeMailer{}) + svc, n, _ := newAnnouncingService(t, newFakeMailer()) inc := &Incident{ID: "inc-1", Title: "Database down", Status: IncidentResolved} svc.AnnounceIncidentUpdate(context.Background(), inc, &IncidentUpdate{IncidentID: "inc-1", Status: IncidentResolved, Message: "Post-mortem published."}) @@ -167,7 +167,7 @@ func TestAnnounceIncidentUpdateAfterResolutionIsAnUpdate(t *testing.T) { func TestNotifySubscribersOutlivesTheRequest(t *testing.T) { pinEdition(t, extension.Pro) - svc, n, _ := newAnnouncingService(t, &fakeMailer{}) + svc, n, _ := newAnnouncingService(t, newFakeMailer()) ctx, cancel := context.WithCancel(context.Background()) cancel() @@ -185,7 +185,7 @@ func TestNotifySubscribersSilentWhileSubscriptionsAreClosed(t *testing.T) { mailer Mailer }{ {"no SMTP server", extension.Pro, nil}, - {"edition without subscribers", extension.Personal, &fakeMailer{}}, + {"edition without subscribers", extension.Personal, newFakeMailer()}, } for _, tc := range cases { t.Run(tc.name, func(t *testing.T) { @@ -207,32 +207,45 @@ type tokenSubscriberStore struct { confirmToken string } -func (s *tokenSubscriberStore) CreateSubscriber(ctx context.Context, sub *StatusSubscriber) (string, error) { +func (s *tokenSubscriberStore) UpsertPendingSubscriber(ctx context.Context, sub *StatusSubscriber) (bool, error) { s.confirmToken = *sub.ConfirmToken - return s.recordingSubscriberStore.CreateSubscriber(ctx, sub) + return s.recordingSubscriberStore.UpsertPendingSubscriber(ctx, sub) } func TestSubscribeEmailsTheConfirmationLink(t *testing.T) { pinEdition(t, extension.Pro) store := &tokenSubscriberStore{} - mailer := &fakeMailer{} + mailer := newFakeMailer() svc := NewSubscriberService(store, mailer, "https://status.example.com/", slog.New(slog.NewTextHandler(io.Discard, nil))) if err := svc.Subscribe(context.Background(), "visitor@example.com"); err != nil { t.Fatalf("subscribe: %v", err) } - mails := mailer.mails() - if len(mails) != 1 || mails[0].to != "visitor@example.com" { - t.Fatalf("mails %+v, want one to the visitor", mails) + mail := mailer.next(t) + if mail.to != "visitor@example.com" { + t.Fatalf("mail to %s, want the visitor", mail.to) } - link := regexp.MustCompile(`https://status\.example\.com/status/confirm\?token=([0-9a-f]{64})`).FindStringSubmatch(mails[0].body) + link := regexp.MustCompile(`https://status\.example\.com/status/confirm\?token=([0-9a-f]{64})`).FindStringSubmatch(mail.body) if link == nil { - t.Fatalf("body %q holds no confirmation link", mails[0].body) + t.Fatalf("body %q holds no confirmation link", mail.body) } if link[1] != store.confirmToken { t.Fatalf("link token %q is not the stored token %q", link[1], store.confirmToken) } + mailer.none(t) +} + +func TestSubscribeSendsNothingToAConfirmedAddress(t *testing.T) { + pinEdition(t, extension.Pro) + store := &recordingSubscriberStore{confirmed: map[string]bool{"visitor@example.com": true}} + mailer := newFakeMailer() + svc := NewSubscriberService(store, mailer, "https://status.example.com", slog.New(slog.NewTextHandler(io.Discard, nil))) + + if err := svc.Subscribe(context.Background(), "visitor@example.com"); err != nil { + t.Fatalf("subscribe: %v", err) + } + mailer.none(t) } func TestSubscribeRefusedWhileSubscriptionsAreClosed(t *testing.T) { diff --git a/internal/status/handler.go b/internal/status/handler.go index ca5a812b..ada1f198 100644 --- a/internal/status/handler.go +++ b/internal/status/handler.go @@ -280,17 +280,15 @@ func (h *Handler) HandleSubscribe(w http.ResponseWriter, r *http.Request) { if err := h.service.subscribers.Subscribe(r.Context(), req.Email); err != nil { h.logger.Error("subscribe failed", "error", err) - switch { - case errors.Is(err, ErrSubscriptionsDisabled): + if errors.Is(err, ErrSubscriptionsDisabled) { writeSubscriptionsUnavailable(w) - case errors.Is(err, ErrConfirmationNotSent): - writeJSONError(w, http.StatusBadGateway, "confirmation_failed", "The confirmation email could not be sent, try again later") - default: - writeJSONError(w, http.StatusInternalServerError, "subscription_failed", "Subscription failed") + return } + writeJSONError(w, http.StatusInternalServerError, "subscription_failed", "Subscription failed") return } + // New, pending or already confirmed: the answer is the same, so it reveals nobody's subscription. w.Header().Set("Content-Type", "application/json") _ = json.NewEncoder(w).Encode(map[string]string{"status": "confirmation_sent"}) } diff --git a/internal/status/handler_test.go b/internal/status/handler_test.go index 04c843ae..640cb168 100644 --- a/internal/status/handler_test.go +++ b/internal/status/handler_test.go @@ -12,7 +12,7 @@ import ( "net/http" "net/http/httptest" "strings" - "sync" + "time" "testing" "github.com/kolapsis/maintenant/internal/extension" @@ -21,18 +21,13 @@ import ( type recordingSubscriberStore struct { SubscriberStore - created []string - deleted []string + created []string + confirmed map[string]bool } -func (s *recordingSubscriberStore) CreateSubscriber(_ context.Context, sub *StatusSubscriber) (string, error) { +func (s *recordingSubscriberStore) UpsertPendingSubscriber(_ context.Context, sub *StatusSubscriber) (bool, error) { s.created = append(s.created, sub.Email) - return "id-" + sub.Email, nil -} - -func (s *recordingSubscriberStore) DeleteSubscriber(_ context.Context, id string) error { - s.deleted = append(s.deleted, id) - return nil + return !s.confirmed[sub.Email], nil } type sentMail struct { @@ -40,25 +35,40 @@ type sentMail struct { } type fakeMailer struct { - mu sync.Mutex - sent []sentMail + sent chan sentMail err error } +func newFakeMailer() *fakeMailer { + return &fakeMailer{sent: make(chan sentMail, 64)} +} + func (m *fakeMailer) Send(_ context.Context, to, subject, body string) error { - m.mu.Lock() - defer m.mu.Unlock() if m.err != nil { return m.err } - m.sent = append(m.sent, sentMail{to: to, subject: subject, body: body}) + m.sent <- sentMail{to: to, subject: subject, body: body} return nil } -func (m *fakeMailer) mails() []sentMail { - m.mu.Lock() - defer m.mu.Unlock() - return append([]sentMail(nil), m.sent...) +func (m *fakeMailer) next(t *testing.T) sentMail { + t.Helper() + select { + case mail := <-m.sent: + return mail + case <-time.After(2 * time.Second): + t.Fatal("no email was sent") + return sentMail{} + } +} + +func (m *fakeMailer) none(t *testing.T) { + t.Helper() + select { + case mail := <-m.sent: + t.Fatalf("unexpected email %q to %s", mail.subject, mail.to) + case <-time.After(100 * time.Millisecond): + } } func pinEdition(t *testing.T, e extension.Edition) { @@ -82,7 +92,7 @@ func newSubscribeHandlerWith(t *testing.T, mailer Mailer) (*Handler, *recordingS func newSubscribeHandler(t *testing.T) (*Handler, *recordingSubscriberStore) { t.Helper() pinEdition(t, extension.Pro) - return newSubscribeHandlerWith(t, &fakeMailer{}) + return newSubscribeHandlerWith(t, newFakeMailer()) } func postSubscribe(t *testing.T, h *Handler, body string) *httptest.ResponseRecorder { @@ -241,8 +251,8 @@ func TestHandleSubscribeRefusedWhileSubscriptionsAreClosed(t *testing.T) { mailer Mailer }{ {"no SMTP server", extension.Pro, nil}, - {"edition without subscribers", extension.Personal, &fakeMailer{}}, - {"community", extension.Community, &fakeMailer{}}, + {"edition without subscribers", extension.Personal, newFakeMailer()}, + {"community", extension.Community, newFakeMailer()}, } for _, tc := range cases { t.Run(tc.name, func(t *testing.T) { @@ -267,20 +277,73 @@ func TestHandleSubscribeRefusedWhileSubscriptionsAreClosed(t *testing.T) { } } -func TestHandleSubscribeDropsTheSubscriberWhenTheConfirmationFails(t *testing.T) { +func TestHandleSubscribeAnswersAlikeWhateverTheAddressState(t *testing.T) { pinEdition(t, extension.Pro) - h, store := newSubscribeHandlerWith(t, &fakeMailer{err: errors.New("connection refused")}) + mailer := newFakeMailer() + h, store := newSubscribeHandlerWith(t, mailer) + store.confirmed = map[string]bool{"known@example.com": true} + + fresh := postSubscribe(t, h, `{"email":"new@example.com"}`) + first := mailer.next(t) + pending := postSubscribe(t, h, `{"email":"new@example.com"}`) + second := mailer.next(t) + confirmed := postSubscribe(t, h, `{"email":"known@example.com"}`) + mailer.none(t) + + for name, rec := range map[string]*httptest.ResponseRecorder{"pending": pending, "confirmed": confirmed} { + if rec.Code != fresh.Code || rec.Body.String() != fresh.Body.String() || rec.Header().Get("Content-Type") != fresh.Header().Get("Content-Type") { + t.Fatalf("%s address: got %d %q, a new one got %d %q", name, rec.Code, rec.Body.String(), fresh.Code, fresh.Body.String()) + } + } + if fresh.Code != http.StatusOK { + t.Fatalf("got %d, want 200", fresh.Code) + } + if first.to != "new@example.com" || second.to != "new@example.com" { + t.Fatalf("confirmations went to %s and %s", first.to, second.to) + } + if first.body == second.body { + t.Fatal("a pending address must get a fresh confirmation link") + } +} - rec := postSubscribe(t, h, `{"email":"ok@example.com"}`) +type stalledMailer struct { + release chan struct{} +} - if rec.Code != http.StatusBadGateway { - t.Fatalf("got %d, want 502", rec.Code) - } - if code := errorCode(t, rec); code != "confirmation_failed" { - t.Fatalf("code %q, want confirmation_failed", code) +func (m *stalledMailer) Send(context.Context, string, string, string) error { + <-m.release + return nil +} + +func TestHandleSubscribeDoesNotWaitForTheMailServer(t *testing.T) { + pinEdition(t, extension.Pro) + mailer := &stalledMailer{release: make(chan struct{})} + t.Cleanup(func() { close(mailer.release) }) + h, _ := newSubscribeHandlerWith(t, mailer) + + done := make(chan int, 1) + go func() { done <- postSubscribe(t, h, `{"email":"ok@example.com"}`).Code }() + + select { + case code := <-done: + if code != http.StatusOK { + t.Fatalf("got %d, want 200", code) + } + case <-time.After(2 * time.Second): + t.Fatal("the answer waited for the mail server, so its delay tells a new address from a confirmed one") } - if len(store.deleted) != 1 || store.deleted[0] != "id-ok@example.com" { - t.Fatalf("deleted %v, want the subscriber just created", store.deleted) +} + +func TestHandleSubscribeAnswersAlikeWhenTheConfirmationFails(t *testing.T) { + pinEdition(t, extension.Pro) + mailer := newFakeMailer() + mailer.err = errors.New("connection refused") + h, _ := newSubscribeHandlerWith(t, mailer) + + rec := postSubscribe(t, h, `{"email":"ok@example.com"}`) + + if rec.Code != http.StatusOK || !strings.Contains(rec.Body.String(), "confirmation_sent") { + t.Fatalf("got %d %q: the answer must not depend on the mail server", rec.Code, rec.Body.String()) } } @@ -309,9 +372,9 @@ func TestStatusAPIReportsWhetherSubscriptionsAreOpen(t *testing.T) { mailer Mailer want bool }{ - {"pro with SMTP", extension.Pro, &fakeMailer{}, true}, + {"pro with SMTP", extension.Pro, newFakeMailer(), true}, {"pro without SMTP", extension.Pro, nil, false}, - {"personal with SMTP", extension.Personal, &fakeMailer{}, false}, + {"personal with SMTP", extension.Personal, newFakeMailer(), false}, } for _, tc := range cases { t.Run(tc.name, func(t *testing.T) { diff --git a/internal/status/store.go b/internal/status/store.go index be63e897..5b3c0ba1 100644 --- a/internal/status/store.go +++ b/internal/status/store.go @@ -38,7 +38,7 @@ type IncidentStore interface { // SubscriberStore defines the persistence interface for email subscribers. type SubscriberStore interface { - CreateSubscriber(ctx context.Context, s *StatusSubscriber) (string, error) + UpsertPendingSubscriber(ctx context.Context, s *StatusSubscriber) (issued bool, err error) GetSubscriberByToken(ctx context.Context, confirmToken string) (*StatusSubscriber, error) GetSubscriberByUnsubToken(ctx context.Context, unsubToken string) (*StatusSubscriber, error) ConfirmSubscriber(ctx context.Context, id string) error diff --git a/internal/status/subscriber.go b/internal/status/subscriber.go index c189b129..d0b8216d 100644 --- a/internal/status/subscriber.go +++ b/internal/status/subscriber.go @@ -19,9 +19,6 @@ import ( // ErrSubscriptionsDisabled is returned when no mailer is configured or the running edition does not open subscribers. var ErrSubscriptionsDisabled = errors.New("email subscriptions are not enabled") -// ErrConfirmationNotSent is returned when the confirmation email could not be delivered; the subscription is dropped. -var ErrConfirmationNotSent = errors.New("confirmation email could not be sent") - // SubscriberService manages email subscriptions for status updates. type SubscriberService struct { store SubscriberStore @@ -54,7 +51,7 @@ func generateToken() (string, error) { return hex.EncodeToString(b), nil } -// Subscribe creates an unconfirmed subscriber and emails it the confirmation link of the double opt-in. +// Subscribe records a pending subscription and, unless the address is already confirmed, emails a fresh confirmation link in the background. func (s *SubscriberService) Subscribe(ctx context.Context, email string) error { if !s.Enabled() { return ErrSubscriptionsDisabled @@ -70,31 +67,28 @@ func (s *SubscriberService) Subscribe(ctx context.Context, email string) error { } expires := time.Now().Add(24 * time.Hour) - sub := &StatusSubscriber{ + issued, err := s.store.UpsertPendingSubscriber(ctx, &StatusSubscriber{ Email: email, - Confirmed: false, ConfirmToken: &confirmToken, ConfirmExpires: &expires, UnsubToken: unsubToken, - } - - id, err := s.store.CreateSubscriber(ctx, sub) + }) if err != nil { - return fmt.Errorf("create subscriber: %w", err) + return fmt.Errorf("record subscriber: %w", err) } + if issued { + go s.sendConfirmation(context.WithoutCancel(ctx), email, confirmToken) + } + return nil +} - confirmURL := fmt.Sprintf("%s/status/confirm?token=%s", s.baseURL, confirmToken) +func (s *SubscriberService) sendConfirmation(ctx context.Context, email, token string) { + confirmURL := fmt.Sprintf("%s/status/confirm?token=%s", s.baseURL, token) body := fmt.Sprintf("Confirm your subscription to status updates by opening this link:\n\n%s\n\n"+ "The link expires in 24 hours. If you did not ask for this, ignore this email.\n", confirmURL) if err := s.mailer.Send(ctx, email, "Confirm your status page subscription", body); err != nil { s.logger.Error("failed to send confirmation email", "error", err, "email", email) - if delErr := s.store.DeleteSubscriber(context.WithoutCancel(ctx), id); delErr != nil { - s.logger.Error("failed to drop unconfirmable subscriber", "error", delErr, "email", email) - } - return fmt.Errorf("%w: %w", ErrConfirmationNotSent, err) } - - return nil } // Confirm validates a confirmation token and activates the subscription. diff --git a/internal/store/subscribers.go b/internal/store/subscribers.go index db32b67b..d5c2e977 100644 --- a/internal/store/subscribers.go +++ b/internal/store/subscribers.go @@ -28,7 +28,8 @@ func NewSubscriberStore(d *DB) *SubscriberStoreImpl { } } -func (s *SubscriberStoreImpl) CreateSubscriber(ctx context.Context, sub *status.StatusSubscriber) (string, error) { +// UpsertPendingSubscriber records sub as an unconfirmed subscription, or hands a still unconfirmed one sub's confirmation token and a fresh 24 h window; issued is false when the address is already confirmed. +func (s *SubscriberStoreImpl) UpsertPendingSubscriber(ctx context.Context, sub *status.StatusSubscriber) (bool, error) { now := time.Now().Unix() var confirmExpires *int64 if sub.ConfirmExpires != nil { @@ -36,17 +37,20 @@ func (s *SubscriberStoreImpl) CreateSubscriber(ctx context.Context, sub *status. confirmExpires = &v } - sub.ID = uid.New() - _, err := s.writer.Exec(ctx, + res, err := s.writer.Exec(ctx, `INSERT INTO status_subscribers (id, email, confirmed, confirm_token, confirm_expires, unsub_token, created_at) - VALUES (?, ?, ?, ?, ?, ?, ?)`, - sub.ID, sub.Email, boolToInt(sub.Confirmed), sub.ConfirmToken, confirmExpires, sub.UnsubToken, now, + VALUES (?, ?, 0, ?, ?, ?, ?) + ON CONFLICT (email) DO UPDATE SET + confirm_token = excluded.confirm_token, + confirm_expires = excluded.confirm_expires, + created_at = excluded.created_at + WHERE status_subscribers.confirmed = 0`, + uid.New(), sub.Email, sub.ConfirmToken, confirmExpires, sub.UnsubToken, now, ) if err != nil { - return "", fmt.Errorf("create subscriber: %w", err) + return false, fmt.Errorf("upsert subscriber: %w", err) } - sub.CreatedAt = time.Unix(now, 0).UTC() - return sub.ID, nil + return res.RowsAffected == 1, nil } func (s *SubscriberStoreImpl) GetSubscriberByToken(ctx context.Context, confirmToken string) (*status.StatusSubscriber, error) { diff --git a/internal/store/subscribers_test.go b/internal/store/subscribers_test.go new file mode 100644 index 00000000..1864a87e --- /dev/null +++ b/internal/store/subscribers_test.go @@ -0,0 +1,87 @@ +// Copyright 2026 Benjamin Touchard (Kolapsis) +// SPDX-License-Identifier: Apache-2.0 + +package store + +import ( + "context" + "testing" + "time" + + "github.com/stretchr/testify/assert" + "github.com/stretchr/testify/require" + + "github.com/kolapsis/maintenant/internal/status" +) + +func pendingSubscription(email, confirmToken, unsubToken string) *status.StatusSubscriber { + expires := time.Now().Add(24 * time.Hour) + return &status.StatusSubscriber{ + Email: email, + ConfirmToken: &confirmToken, + ConfirmExpires: &expires, + UnsubToken: unsubToken, + } +} + +func TestUpsertPendingSubscriber_NewPendingConfirmed(t *testing.T) { + db := openTestDB(t) + s := NewSubscriberStore(db) + ctx := context.Background() + + issued, err := s.UpsertPendingSubscriber(ctx, pendingSubscription("visitor@example.com", "confirm-1", "unsub-1")) + require.NoError(t, err) + assert.True(t, issued, "a new address gets a confirmation token") + + issued, err = s.UpsertPendingSubscriber(ctx, pendingSubscription("visitor@example.com", "confirm-2", "unsub-2")) + require.NoError(t, err) + assert.True(t, issued, "a pending address gets a fresh confirmation token") + + stale, err := s.GetSubscriberByToken(ctx, "confirm-1") + require.NoError(t, err) + assert.Nil(t, stale, "the previous confirmation link stops working") + pending, err := s.GetSubscriberByToken(ctx, "confirm-2") + require.NoError(t, err) + require.NotNil(t, pending) + assert.Equal(t, "unsub-1", pending.UnsubToken, "the unsubscribe link of the address does not change") + assert.False(t, pending.Confirmed) + + require.NoError(t, s.ConfirmSubscriber(ctx, pending.ID)) + + issued, err = s.UpsertPendingSubscriber(ctx, pendingSubscription("visitor@example.com", "confirm-3", "unsub-3")) + require.NoError(t, err) + assert.False(t, issued, "a confirmed address gets no confirmation token") + + again, err := s.GetSubscriberByToken(ctx, "confirm-3") + require.NoError(t, err) + assert.Nil(t, again) + confirmed, err := s.GetSubscriberByUnsubToken(ctx, "unsub-1") + require.NoError(t, err) + require.NotNil(t, confirmed) + assert.True(t, confirmed.Confirmed, "subscribing again never unconfirms an address") + assert.Nil(t, confirmed.ConfirmToken) + + stats, err := s.GetSubscriberStats(ctx) + require.NoError(t, err) + assert.Equal(t, 1, stats.Total) + assert.Equal(t, 1, stats.Confirmed) +} + +func TestUpsertPendingSubscriber_RefreshKeepsAPendingAddressAlive(t *testing.T) { + db := openTestDB(t) + s := NewSubscriberStore(db) + ctx := context.Background() + + _, err := s.UpsertPendingSubscriber(ctx, pendingSubscription("visitor@example.com", "confirm-1", "unsub-1")) + require.NoError(t, err) + _, err = db.Writer().Exec(ctx, `UPDATE status_subscribers SET created_at = ? WHERE email = ?`, + time.Now().Add(-25*time.Hour).Unix(), "visitor@example.com") + require.NoError(t, err) + + _, err = s.UpsertPendingSubscriber(ctx, pendingSubscription("visitor@example.com", "confirm-2", "unsub-2")) + require.NoError(t, err) + + deleted, err := s.CleanExpiredUnconfirmed(ctx) + require.NoError(t, err) + assert.Zero(t, deleted, "a fresh confirmation link must not be swept by the 24 h cleanup") +} From 937461402106c10c69d41a56fcfca008f2d46994 Mon Sep 17 00:00:00 2001 From: Benjamin Date: Wed, 30 Sep 2026 19:14:58 +0200 Subject: [PATCH 11/54] fix: persist daily uptime, honour ignore and Swarm service labels, add Kubernetes service insights Daily uptime is now rolled up into endpoint_, heartbeat_ and container_uptime_daily (migration 35, both engines) before the raw rows are purged, and kept 365 days; the current day is still computed live with portable SQL. Heartbeat days count completion pings only, exit code 0 as success. Transition purge keeps each container's latest transition. A container marked maintenant.ignore raises no alert and declares no endpoint or certificate, locally, through agents and from Kubernetes annotations; label-derived fields are refreshed on reconcile. Swarm task containers read their service labels (deploy.labels) under their own and are grouped by stack; the unused task-to-container mapping is removed. LoadBalancer and NodePort Services exposing a local workload raise security insights; missing_network_policy is removed. --- .../src/components/SecurityInsightList.vue | 2 - frontend/src/services/securityApi.ts | 1 - internal/agent/collector.go | 2 +- internal/api/v1/runtime_status_test.go | 2 +- internal/app/app.go | 5 +- internal/app/ignore_test.go | 49 ++ internal/app/lifecycle.go | 42 +- internal/app/security.go | 85 ++- internal/app/security_kubernetes_test.go | 73 ++ internal/certificate/labels.go | 7 +- internal/certificate/service_test.go | 9 + internal/container/agent_event.go | 9 +- internal/container/ignore_test.go | 162 ++++ internal/container/labels.go | 46 ++ internal/container/service.go | 21 +- internal/docker/client.go | 9 +- internal/docker/discovery.go | 34 +- internal/docker/events.go | 2 +- internal/docker/service_labels.go | 77 ++ internal/docker/service_labels_test.go | 146 ++++ internal/endpoint/labels.go | 8 +- internal/endpoint/service_test.go | 36 + internal/kubernetes/services.go | 166 +++++ internal/kubernetes/services_test.go | 107 +++ internal/resource/agent_event.go | 2 +- internal/resource/agent_event_test.go | 15 + internal/security/analyzer.go | 61 ++ internal/security/analyzer_kubernetes_test.go | 56 ++ internal/security/model.go | 1 - internal/store/containers.go | 8 +- internal/store/copy.go | 3 +- .../postgres/35_uptime_daily.down.sql | 3 + .../postgres/35_uptime_daily.up.sql | 26 + .../sqlite/35_uptime_daily.down.sql | 3 + .../migrations/sqlite/35_uptime_daily.up.sql | 26 + internal/store/migrations_postgres_test.go | 9 +- internal/store/retention.go | 95 ++- internal/store/transform.go | 9 +- internal/store/uptime_daily.go | 526 +++++++++---- internal/store/uptime_daily_test.go | 701 +++++++++--------- internal/store/uuid_schema.sql | 27 + internal/swarm/labels.go | 56 +- internal/swarm/service.go | 128 +--- internal/swarm/snapshot.go | 2 +- 44 files changed, 2057 insertions(+), 800 deletions(-) create mode 100644 internal/app/ignore_test.go create mode 100644 internal/app/security_kubernetes_test.go create mode 100644 internal/container/ignore_test.go create mode 100644 internal/container/labels.go create mode 100644 internal/docker/service_labels.go create mode 100644 internal/docker/service_labels_test.go create mode 100644 internal/kubernetes/services.go create mode 100644 internal/kubernetes/services_test.go create mode 100644 internal/security/analyzer_kubernetes_test.go create mode 100644 internal/store/migrations/postgres/35_uptime_daily.down.sql create mode 100644 internal/store/migrations/postgres/35_uptime_daily.up.sql create mode 100644 internal/store/migrations/sqlite/35_uptime_daily.down.sql create mode 100644 internal/store/migrations/sqlite/35_uptime_daily.up.sql diff --git a/frontend/src/components/SecurityInsightList.vue b/frontend/src/components/SecurityInsightList.vue index 9f6a3a35..6b0cab42 100644 --- a/frontend/src/components/SecurityInsightList.vue +++ b/frontend/src/components/SecurityInsightList.vue @@ -45,8 +45,6 @@ function insightIcon(type: string) { case 'service_load_balancer': case 'service_node_port': return Server - case 'missing_network_policy': - return ShieldAlert default: return ShieldAlert } diff --git a/frontend/src/services/securityApi.ts b/frontend/src/services/securityApi.ts index 382a7346..dd786bc4 100644 --- a/frontend/src/services/securityApi.ts +++ b/frontend/src/services/securityApi.ts @@ -12,7 +12,6 @@ export type InsightType = | 'host_network_mode' | 'service_load_balancer' | 'service_node_port' - | 'missing_network_policy' export interface SecurityInsight { type: InsightType diff --git a/internal/agent/collector.go b/internal/agent/collector.go index c98b7c80..af8f4fce 100644 --- a/internal/agent/collector.go +++ b/internal/agent/collector.go @@ -285,7 +285,7 @@ func collectResourceSnapshots(ctx context.Context, id *Identity, rt runtime.Runt } for _, c := range containers { - if c.State != cmodel.StateRunning { + if c.State != cmodel.StateRunning || c.IsIgnored { continue } raw, err := rt.StatsSnapshot(ctx, c.ExternalID) diff --git a/internal/api/v1/runtime_status_test.go b/internal/api/v1/runtime_status_test.go index 795af661..b130b7f3 100644 --- a/internal/api/v1/runtime_status_test.go +++ b/internal/api/v1/runtime_status_test.go @@ -65,7 +65,7 @@ func discoveryWithServices(t *testing.T, n int) *swarm.ServiceDiscovery { svcs[i] = dockerswarm.Service{ID: "svc" + string(rune('a'+i))} } disc := swarm.NewServiceDiscovery(fakeServiceClient{services: svcs}, testLogger()) - _, _, err := disc.DiscoverAll(context.Background()) + _, err := disc.DiscoverAll(context.Background()) require.NoError(t, err) return disc } diff --git a/internal/app/app.go b/internal/app/app.go index 29029368..f02b8ca8 100644 --- a/internal/app/app.go +++ b/internal/app/app.go @@ -97,6 +97,7 @@ type App struct { hbStore *store.HeartbeatStore certStore *store.CertificateStore resStore *store.ResourceStore + uptimeStore *store.UptimeDailyStore agentStore *store.AgentStore agentSessions extpoint.AgentSessions serveAgents func(ctx context.Context, cfg extpoint.GRPCConfig) error @@ -606,7 +607,7 @@ func New(cfg Config, logger *slog.Logger, opts ...Option) (*App, error) { a.wireAgentLifecycleAlerts() // --- Router --- - uptimeDailyStore := store.NewUptimeDailyStore(db) + a.uptimeStore = store.NewUptimeDailyStore(db) a.router = v1.NewRouter(v1.HandlerDeps{ // Core services Broker: a.broker, @@ -639,7 +640,7 @@ func New(cfg Config, logger *slog.Logger, opts ...Option) (*App, error) { // Webhooks WebhookStore: webhookStore, // UI extras - UptimeDaily: uptimeDailyStore, + UptimeDaily: a.uptimeStore, LogStreamer: rt, ResourceTopSvc: a.resourceSvc, SparklineFetcher: epStore, diff --git a/internal/app/ignore_test.go b/internal/app/ignore_test.go new file mode 100644 index 00000000..97318da5 --- /dev/null +++ b/internal/app/ignore_test.go @@ -0,0 +1,49 @@ +// Copyright 2026 Benjamin Touchard (Kolapsis) +// SPDX-License-Identifier: Apache-2.0 + +package app + +import ( + "testing" + "time" + + "github.com/stretchr/testify/assert" + + "github.com/kolapsis/maintenant/internal/container" + "github.com/kolapsis/maintenant/internal/docker" + "github.com/kolapsis/maintenant/internal/runtime" +) + +func TestDockerInsights_IgnoredContainerHasNone(t *testing.T) { + cfg := &docker.SecurityConfig{ + Privileged: true, + PortBindings: []docker.PortBindingInfo{{HostIP: "0.0.0.0", HostPort: "5432", ContainerPort: 5432, Protocol: "tcp"}}, + } + + monitored := &container.Container{ID: "c1", Name: "db"} + assert.NotEmpty(t, dockerInsights(monitored, cfg, time.Now())) + + ignored := &container.Container{ID: "c1", Name: "db", IsIgnored: true} + assert.Empty(t, dockerInsights(ignored, cfg, time.Now())) +} + +func TestSwarmTaskFailure(t *testing.T) { + task := map[string]string{ + "com.docker.swarm.service.id": "svc1", + "com.docker.swarm.service.name": "prod_web", + } + id, name, ok := swarmTaskFailure(runtime.RuntimeEvent{Action: "die", Labels: task}) + assert.True(t, ok) + assert.Equal(t, "svc1", id) + assert.Equal(t, "prod_web", name) + + _, _, ok = swarmTaskFailure(runtime.RuntimeEvent{Action: "start", Labels: task}) + assert.False(t, ok) + + _, _, ok = swarmTaskFailure(runtime.RuntimeEvent{Action: "die", Labels: map[string]string{"name": "standalone"}}) + assert.False(t, ok) + + ignored := map[string]string{"com.docker.swarm.service.id": "svc1", "maintenant.ignore": "true"} + _, _, ok = swarmTaskFailure(runtime.RuntimeEvent{Action: "die", Labels: ignored}) + assert.False(t, ok, "an ignored task feeds no crash-loop alert") +} diff --git a/internal/app/lifecycle.go b/internal/app/lifecycle.go index fed7f4ff..a686ec10 100644 --- a/internal/app/lifecycle.go +++ b/internal/app/lifecycle.go @@ -15,7 +15,6 @@ import ( "github.com/kolapsis/maintenant/internal/event" "github.com/kolapsis/maintenant/internal/kubernetes" "github.com/kolapsis/maintenant/internal/runtime" - "github.com/kolapsis/maintenant/internal/security" "github.com/kolapsis/maintenant/internal/store" "github.com/kolapsis/maintenant/internal/swarm" "github.com/kolapsis/maintenant/internal/uid" @@ -92,7 +91,7 @@ func (a *App) reconcile(ctx context.Context) { // Swarm service discovery on startup. if a.swarmDiscovery != nil { a.logger.Info("running Swarm service discovery") - _, services, err := a.swarmDiscovery.DiscoverAll(ctx) + services, err := a.swarmDiscovery.DiscoverAll(ctx) if err != nil { a.logger.Error("Swarm service discovery failed", "error", err) } else { @@ -130,21 +129,7 @@ func (a *App) reconcile(ctx context.Context) { dbC := dbByExtID[r.Container.ExternalID] if r.SecurityConfig != nil && dbC != nil && dbC.ID != "" { - bindings := make([]security.PortBinding, 0, len(r.SecurityConfig.PortBindings)) - for _, pb := range r.SecurityConfig.PortBindings { - bindings = append(bindings, security.PortBinding{ - HostIP: pb.HostIP, - HostPort: pb.HostPort, - Port: pb.ContainerPort, - Protocol: pb.Protocol, - }) - } - insights := security.AnalyzeDocker(dbC.ID, dbC.Name, security.DockerSecurityConfig{ - Privileged: r.SecurityConfig.Privileged, - NetworkMode: r.SecurityConfig.NetworkMode, - Bindings: bindings, - }, now) - a.securitySvc.UpdateContainer(dbC.ID, dbC.Name, insights) + a.securitySvc.UpdateContainer(dbC.ID, dbC.Name, dockerInsights(dbC, r.SecurityConfig, now)) } } if swept := a.endpointSvc.SweepOrphanedLabelEndpoints(ctx, seen); swept > 0 { @@ -204,7 +189,7 @@ func (a *App) startEventStream(ctx context.Context) <-chan struct{} { name = name[1:] } a.endpointSvc.HandleContainerStart(ctx, name, evt.ExternalID, evt.Labels, - evt.Labels["com.docker.compose.project"], + container.OrchestrationGroupFromLabels(evt.Labels), evt.Labels["com.docker.compose.service"]) a.certSvc.SyncFromLabels(ctx, evt.ExternalID, evt.Labels) @@ -215,9 +200,8 @@ func (a *App) startEventStream(ctx context.Context) <-chan struct{} { a.endpointSvc.HandleContainerStop(ctx, evt.ExternalID) // Feed Swarm task failures to crash-loop detector (Pro). - if evt.Action == "die" && a.swarmCrashLoop != nil { - if svcID, ok := evt.Labels["com.docker.swarm.service.id"]; ok && svcID != "" { - svcName := evt.Labels["com.docker.swarm.service.name"] + if a.swarmCrashLoop != nil { + if svcID, svcName, ok := swarmTaskFailure(evt); ok { a.swarmCrashLoop.RecordFailure(svcID, svcName, evt.ErrorDetail) // Emit task_failed SSE event. @@ -243,6 +227,16 @@ func (a *App) startEventStream(ctx context.Context) <-chan struct{} { return done } +// swarmTaskFailure returns the service of a Swarm task that died, unless the +// task is ignored. +func swarmTaskFailure(evt runtime.RuntimeEvent) (serviceID, serviceName string, ok bool) { + serviceID = evt.Labels["com.docker.swarm.service.id"] + if evt.Action != "die" || serviceID == "" || container.IgnoredByLabels(evt.Labels) { + return "", "", false + } + return serviceID, evt.Labels["com.docker.swarm.service.name"], true +} + // startNodeRefresh runs periodic Swarm node reconciliation (Pro, 60s). func (a *App) startNodeRefresh(ctx context.Context) { ticker := time.NewTicker(60 * time.Second) @@ -304,6 +298,9 @@ func (a *App) startKubernetesReconcile(ctx context.Context, src kubernetes.Snaps if err := a.k8sIngest.Reconcile(ctx, uid.LocalAgent, snap); err != nil { a.logger.Warn("local kubernetes reconcile: store failed", "error", err) } + if exposures, ok := a.rt.(serviceExposureSource); ok { + ScanKubernetesSecurity(ctx, exposures, a.containerSvc, a.securitySvc, a.logger) + } } reconcile() ticker := time.NewTicker(localTopologyReconcileInterval) @@ -363,6 +360,7 @@ func (a *App) startRetentionCleanup(ctx context.Context) { HeartbeatStore: a.hbStore, CertificateStore: a.certStore, ResourceStore: a.resStore, + UptimeStore: a.uptimeStore, Config: store.RetentionConfig{ Snapshots: a.cfg.Retention.Snapshots, Interval: a.cfg.Retention.Interval, @@ -465,7 +463,7 @@ func (a *App) startSwarmRecheck(ctx context.Context) { a.swarmEvents = swarm.NewEventProcessor(a.swarmDiscovery, a.logger) // Run initial discovery. - _, services, err := a.swarmDiscovery.DiscoverAll(ctx) + services, err := a.swarmDiscovery.DiscoverAll(ctx) if err != nil { a.logger.Error("initial Swarm discovery after activation failed", "error", err) } else { diff --git a/internal/app/security.go b/internal/app/security.go index 45717c9c..690ee7c0 100644 --- a/internal/app/security.go +++ b/internal/app/security.go @@ -11,7 +11,9 @@ import ( "github.com/kolapsis/maintenant/internal/alert" "github.com/kolapsis/maintenant/internal/container" "github.com/kolapsis/maintenant/internal/docker" + "github.com/kolapsis/maintenant/internal/kubernetes" "github.com/kolapsis/maintenant/internal/security" + "github.com/kolapsis/maintenant/internal/uid" ) // ScanContainerSecurity inspects a single container and updates its security insights. @@ -40,23 +42,78 @@ func ScanContainerSecurity(ctx context.Context, dr *docker.Runtime, containerSvc if c == nil { return } + secSvc.UpdateContainer(c.ID, c.Name, dockerInsights(c, r.SecurityConfig, now)) + return + } +} + +// dockerInsights analyses a container's Docker configuration; an ignored +// container has none. +func dockerInsights(c *container.Container, cfg *docker.SecurityConfig, now time.Time) []security.Insight { + if c.IsIgnored { + return nil + } + bindings := make([]security.PortBinding, 0, len(cfg.PortBindings)) + for _, pb := range cfg.PortBindings { + bindings = append(bindings, security.PortBinding{ + HostIP: pb.HostIP, + HostPort: pb.HostPort, + Port: pb.ContainerPort, + Protocol: pb.Protocol, + }) + } + return security.AnalyzeDocker(c.ID, c.Name, security.DockerSecurityConfig{ + Privileged: cfg.Privileged, + NetworkMode: cfg.NetworkMode, + Bindings: bindings, + }, now) +} + +// serviceExposureSource lists the ports Kubernetes Services open outside the cluster. +type serviceExposureSource interface { + ListServiceExposures(ctx context.Context) ([]kubernetes.ServiceExposure, error) +} - bindings := make([]security.PortBinding, 0, len(r.SecurityConfig.PortBindings)) - for _, pb := range r.SecurityConfig.PortBindings { - bindings = append(bindings, security.PortBinding{ - HostIP: pb.HostIP, - HostPort: pb.HostPort, - Port: pb.ContainerPort, - Protocol: pb.Protocol, - }) +type containerLister interface { + ListContainers(ctx context.Context, opts container.ListContainersOpts) ([]*container.Container, error) +} + +// ScanKubernetesSecurity refreshes the insights of the local cluster's +// workloads from the LoadBalancer and NodePort Services that expose them. +func ScanKubernetesSecurity(ctx context.Context, src serviceExposureSource, containers containerLister, secSvc *security.Service, logger *slog.Logger) { + exposures, err := src.ListServiceExposures(ctx) + if err != nil { + logger.Warn("security: kubernetes services not analysed", "error", err) + return + } + byWorkload := make(map[string][]security.ServicePort, len(exposures)) + for _, e := range exposures { + byWorkload[e.WorkloadID] = append(byWorkload[e.WorkloadID], security.ServicePort{ + Service: e.Service, + Type: e.ServiceType, + Port: e.Port, + TargetPort: e.TargetPort, + NodePort: e.NodePort, + Protocol: e.Protocol, + }) + } + + local := uid.LocalAgent + workloads, err := containers.ListContainers(ctx, container.ListContainersOpts{IncludeIgnored: true, AgentFilter: &local}) + if err != nil { + logger.Warn("security: kubernetes workloads not listed", "error", err) + return + } + now := time.Now() + for _, c := range workloads { + if c.RuntimeType != "kubernetes" { + continue + } + var insights []security.Insight + if !c.IsIgnored { + insights = security.AnalyzeKubernetes(c.ID, c.Name, byWorkload[c.ExternalID], now) } - insights := security.AnalyzeDocker(c.ID, c.Name, security.DockerSecurityConfig{ - Privileged: r.SecurityConfig.Privileged, - NetworkMode: r.SecurityConfig.NetworkMode, - Bindings: bindings, - }, now) secSvc.UpdateContainer(c.ID, c.Name, insights) - return } } diff --git a/internal/app/security_kubernetes_test.go b/internal/app/security_kubernetes_test.go new file mode 100644 index 00000000..c82fc53f --- /dev/null +++ b/internal/app/security_kubernetes_test.go @@ -0,0 +1,73 @@ +// Copyright 2026 Benjamin Touchard (Kolapsis) +// SPDX-License-Identifier: Apache-2.0 + +package app + +import ( + "context" + "io" + "log/slog" + "testing" + + "github.com/stretchr/testify/assert" + "github.com/stretchr/testify/require" + + "github.com/kolapsis/maintenant/internal/container" + "github.com/kolapsis/maintenant/internal/kubernetes" + "github.com/kolapsis/maintenant/internal/security" + "github.com/kolapsis/maintenant/internal/uid" +) + +type fakeExposures []kubernetes.ServiceExposure + +func (f fakeExposures) ListServiceExposures(context.Context) ([]kubernetes.ServiceExposure, error) { + return f, nil +} + +type fakeWorkloads []*container.Container + +func (f fakeWorkloads) ListContainers(_ context.Context, opts container.ListContainersOpts) ([]*container.Container, error) { + if !opts.IncludeIgnored || opts.AgentFilter == nil || *opts.AgentFilter != uid.LocalAgent { + return nil, nil + } + return f, nil +} + +type securityAlert struct { + containerID string + recover bool +} + +func TestScanKubernetesSecurity(t *testing.T) { + logger := slog.New(slog.NewTextHandler(io.Discard, nil)) + var alerts []securityAlert + secSvc := security.NewService(security.Deps{ + Logger: logger, + AlertCallback: func(id string, _ string, _ []security.Insight, isRecover bool) { + alerts = append(alerts, securityAlert{id, isRecover}) + }, + }) + + workloads := fakeWorkloads{ + {ID: "w-web", ExternalID: "shop/Deployment/web", Name: "web", RuntimeType: "kubernetes"}, + {ID: "w-db", ExternalID: "shop/StatefulSet/db", Name: "db", RuntimeType: "kubernetes", IsIgnored: true}, + {ID: "w-local", ExternalID: "abc", Name: "sidecar", RuntimeType: "docker"}, + } + exposed := fakeExposures{ + {WorkloadID: "shop/Deployment/web", Service: "shop/web-lb", ServiceType: "LoadBalancer", Port: 443, TargetPort: 8443, Protocol: "tcp"}, + {WorkloadID: "shop/StatefulSet/db", Service: "shop/db-np", ServiceType: "NodePort", Port: 5432, TargetPort: 5432, NodePort: 31432, Protocol: "tcp"}, + } + + ScanKubernetesSecurity(context.Background(), exposed, workloads, secSvc, logger) + + web := secSvc.GetContainerInsights("w-web") + require.Equal(t, 1, web.Count) + assert.Equal(t, security.ServiceLoadBalancer, web.Insights[0].Type) + assert.Zero(t, secSvc.GetContainerInsights("w-db").Count, "an ignored workload raises nothing") + assert.Equal(t, []securityAlert{{"w-web", false}}, alerts, "the dangerous_configuration circuit is fed") + + ScanKubernetesSecurity(context.Background(), fakeExposures{}, workloads, secSvc, logger) + + assert.Zero(t, secSvc.GetContainerInsights("w-web").Count) + assert.Equal(t, securityAlert{"w-web", true}, alerts[len(alerts)-1], "a Service removed resolves its alert") +} diff --git a/internal/certificate/labels.go b/internal/certificate/labels.go index b6be9d42..b6da4612 100644 --- a/internal/certificate/labels.go +++ b/internal/certificate/labels.go @@ -7,6 +7,8 @@ import ( "net" "strconv" "strings" + + "github.com/kolapsis/maintenant/internal/container" ) const tlsLabel = "maintenant.tls.certificates" @@ -19,8 +21,11 @@ type ParsedCertLabel struct { // ParseCertificateLabels extracts certificate monitoring targets from container labels. // The label format is: maintenant.tls.certificates=host1,host2:8443,host3 -// Hostnames without a port default to 443. +// Hostnames without a port default to 443. An ignored container declares none. func ParseCertificateLabels(labels map[string]string) []ParsedCertLabel { + if container.IgnoredByLabels(labels) { + return nil + } raw, ok := labels[tlsLabel] if !ok { return nil diff --git a/internal/certificate/service_test.go b/internal/certificate/service_test.go index e03c239f..5b226bfe 100644 --- a/internal/certificate/service_test.go +++ b/internal/certificate/service_test.go @@ -280,6 +280,15 @@ func TestParseCertificateLabels_StripsSchemeAndPath(t *testing.T) { assert.Equal(t, 443, parsed[0].Port) } +func TestParseCertificateLabels_IgnoredContainerDeclaresNone(t *testing.T) { + parsed := ParseCertificateLabels(map[string]string{ + "maintenant.ignore": "true", + "maintenant.tls.certificates": "example.com", + }) + + assert.Empty(t, parsed) +} + // --------------------------------------------------------------------------- // Mock store for quota testing // --------------------------------------------------------------------------- diff --git a/internal/container/agent_event.go b/internal/container/agent_event.go index c3cc90b7..ddc9b903 100644 --- a/internal/container/agent_event.go +++ b/internal/container/agent_event.go @@ -98,6 +98,9 @@ func (s *Service) refreshAgentContainer(ctx context.Context, c *Container, ev *a if labels := ev.GetLabels(); len(labels) > 0 && c.ApplyImageLabels(labels) { dirty = true } + if labels := ev.GetLabels(); len(labels) > 0 && c.adoptLabelFields(agentLabelFields(labels)) { + dirty = true + } if ev.GetHasHealthCheck() && !c.HasHealthCheck { c.HasHealthCheck = true dirty = true @@ -186,7 +189,7 @@ func (s *Service) insertAgentContainer(ctx context.Context, agentID string, ev * Name: ev.GetName(), Image: ev.GetImage(), State: state, - OrchestrationGroup: labels[labelComposeProject], + OrchestrationGroup: OrchestrationGroupFromLabels(labels), OrchestrationUnit: labels[labelComposeService], ComposeWorkingDir: labels[labelComposeWorkingDir], RuntimeType: s.resolveAgentRuntime(ctx, agentID), @@ -238,7 +241,9 @@ func (s *Service) insertAgentContainer(ctx context.Context, agentID string, ev * s.logger.Info("agent event: container discovered", "external_id", shortID(externalID), "name", c.Name, "agent_id", agentID, "state", string(state)) - s.emitEvent(event.ContainerDiscovered, c) + if !c.IsIgnored { + s.emitEvent(event.ContainerDiscovered, c) + } return nil } diff --git a/internal/container/ignore_test.go b/internal/container/ignore_test.go new file mode 100644 index 00000000..f91d86aa --- /dev/null +++ b/internal/container/ignore_test.go @@ -0,0 +1,162 @@ +// Copyright 2026 Benjamin Touchard (Kolapsis) +// SPDX-License-Identifier: Apache-2.0 + +package container + +import ( + "context" + "testing" + "time" + + "github.com/stretchr/testify/assert" + "github.com/stretchr/testify/require" + + "github.com/kolapsis/maintenant/internal/agentevent" + "github.com/kolapsis/maintenant/internal/agentpb" +) + +func TestIgnoredByLabels(t *testing.T) { + assert.True(t, IgnoredByLabels(map[string]string{"maintenant.ignore": "true"})) + assert.True(t, IgnoredByLabels(map[string]string{"maintenant.ignore": "1"})) + assert.False(t, IgnoredByLabels(map[string]string{"maintenant.ignore": "false"})) + assert.False(t, IgnoredByLabels(nil)) +} + +func TestOrchestrationGroupFromLabels(t *testing.T) { + assert.Equal(t, "shop", OrchestrationGroupFromLabels(map[string]string{ + "com.docker.compose.project": "shop", "com.docker.stack.namespace": "prod", + })) + assert.Equal(t, "prod", OrchestrationGroupFromLabels(map[string]string{"com.docker.stack.namespace": "prod"})) + assert.Empty(t, OrchestrationGroupFromLabels(nil)) +} + +func TestService_IgnoredContainerRaisesNothing(t *testing.T) { + store := newSvcStore() + c := makeTestContainer(extID("ignored"), StateExited) + c.ID = "ign-1" + c.IsIgnored = true + healthy := HealthHealthy + c.HealthStatus = &healthy + store.seed(c) + + checker := &mockRestartChecker{result: map[string]interface{}{"restarts": 10}} + var emitted []string + svc := newTestService(store, func(d *Deps) { + d.RestartChecker = checker + d.EventCallback = func(eventType string, _ interface{}) { emitted = append(emitted, eventType) } + }) + + svc.ProcessEvent(context.Background(), makeTestEvent("start", c.ExternalID)) + unhealthy := makeTestEvent("health_status", c.ExternalID) + unhealthy.HealthStatus = string(HealthUnhealthy) + svc.ProcessEvent(context.Background(), unhealthy) + + assert.Empty(t, emitted, "an ignored container must not feed the alert pipeline") + assert.Zero(t, checker.calls, "no restart check for an ignored container") +} + +// Swarm service labels and Kubernetes annotations change on a container that +// is already stored; the next reconcile must pick them up. +func TestService_Reconcile_AdoptsLabelFields(t *testing.T) { + store := newSvcStore() + c := makeTestContainer(extID("stack"), StateRunning) + c.ID = "stk-1" + c.AlertSeverity = SeverityWarning + c.RestartThreshold = 3 + store.seed(c) + + discovered := makeTestContainer(c.ExternalID, StateExited) + discovered.IsIgnored = true + discovered.OrchestrationGroup = "prod" + discovered.CustomGroup = "payments" + discovered.AlertSeverity = SeverityCritical + discovered.RestartThreshold = 7 + + var emitted []string + discoverer := &mockDiscoverer{containers: []*Container{discovered}} + svc := newTestService(store, func(d *Deps) { + d.Discoverer = discoverer + d.EventCallback = func(eventType string, _ interface{}) { emitted = append(emitted, eventType) } + }) + + require.NoError(t, svc.Reconcile(context.Background(), discoverer)) + + got, err := store.GetContainerByID(context.Background(), c.ID) + require.NoError(t, err) + assert.True(t, got.IsIgnored) + assert.Equal(t, "prod", got.OrchestrationGroup) + assert.Equal(t, "payments", got.CustomGroup) + assert.Equal(t, SeverityCritical, got.AlertSeverity) + assert.Equal(t, 7, got.RestartThreshold) + assert.NotContains(t, emitted, "container.state_changed", "an ignored container announces no state change") +} + +func TestService_Reconcile_IgnoredNewContainerIsNotAnnounced(t *testing.T) { + store := newSvcStore() + discovered := makeTestContainer(extID("newign"), StateRunning) + discovered.IsIgnored = true + + var emitted []string + discoverer := &mockDiscoverer{containers: []*Container{discovered}} + svc := newTestService(store, func(d *Deps) { + d.Discoverer = discoverer + d.EventCallback = func(eventType string, _ interface{}) { emitted = append(emitted, eventType) } + }) + + require.NoError(t, svc.Reconcile(context.Background(), discoverer)) + + assert.NotContains(t, emitted, "container.discovered") +} + +func TestHandleAgentEvent_AdoptsLabelFields(t *testing.T) { + store := newSvcStore() + id := extID("agentstack") + c := makeTestContainer(id, StateRunning) + c.ID = "ag-1" + c.AgentID = "agent-1" + c.AlertSeverity = SeverityWarning + c.RestartThreshold = 3 + store.seed(c) + + var emitted []string + svc := newTestService(store, func(d *Deps) { + d.EventCallback = func(eventType string, _ interface{}) { emitted = append(emitted, eventType) } + }) + + ev := &agentpb.ContainerEvent{ + ContainerId: id, + Name: c.Name, + State: agentpb.ContainerState_CONTAINER_STATE_EXITED, + Labels: map[string]string{ + "maintenant.ignore": "true", + "com.docker.stack.namespace": "prod", + }, + } + require.NoError(t, svc.HandleAgentEvent(context.Background(), "agent-1", ev, agentevent.Meta{ObservedAt: time.Now()})) + + got, err := store.GetContainerByExternalID(context.Background(), "agent-1", id) + require.NoError(t, err) + assert.True(t, got.IsIgnored) + assert.Equal(t, "prod", got.OrchestrationGroup) + assert.Equal(t, StateRunning, got.State, "the exit of an ignored container is not processed") + assert.Empty(t, emitted) +} + +func TestHandleAgentEvent_NewContainerGroupedByStack(t *testing.T) { + store := newSvcStore() + id := extID("agentnew") + svc := newTestService(store) + + ev := &agentpb.ContainerEvent{ + ContainerId: id, + Name: "prod_web.1.abc", + State: agentpb.ContainerState_CONTAINER_STATE_RUNNING, + Labels: map[string]string{"com.docker.stack.namespace": "prod"}, + } + require.NoError(t, svc.HandleAgentEvent(context.Background(), "agent-1", ev, agentevent.Meta{ObservedAt: time.Now()})) + + got, err := store.GetContainerByExternalID(context.Background(), "agent-1", id) + require.NoError(t, err) + require.NotNil(t, got) + assert.Equal(t, "prod", got.OrchestrationGroup) +} diff --git a/internal/container/labels.go b/internal/container/labels.go new file mode 100644 index 00000000..50b2a391 --- /dev/null +++ b/internal/container/labels.go @@ -0,0 +1,46 @@ +// Copyright 2026 Benjamin Touchard (Kolapsis) +// SPDX-License-Identifier: Apache-2.0 + +package container + +const labelStackNamespace = "com.docker.stack.namespace" + +// IgnoredByLabels reports whether labels mark a container with maintenant.ignore. +func IgnoredByLabels(labels map[string]string) bool { + v := labels[labelPBIgnore] + return v == "true" || v == "1" +} + +// OrchestrationGroupFromLabels returns the Compose project of a container, or +// its Swarm stack when Compose did not start it. +func OrchestrationGroupFromLabels(labels map[string]string) string { + if project := labels[labelComposeProject]; project != "" { + return project + } + return labels[labelStackNamespace] +} + +// adoptLabelFields copies from d the fields read from labels or annotations and +// reports whether any changed. +func (c *Container) adoptLabelFields(d *Container) bool { + changed := c.IsIgnored != d.IsIgnored || c.CustomGroup != d.CustomGroup || + c.AlertSeverity != d.AlertSeverity || c.RestartThreshold != d.RestartThreshold || + c.OrchestrationGroup != d.OrchestrationGroup + c.IsIgnored = d.IsIgnored + c.CustomGroup = d.CustomGroup + c.AlertSeverity = d.AlertSeverity + c.RestartThreshold = d.RestartThreshold + c.OrchestrationGroup = d.OrchestrationGroup + return changed +} + +// agentLabelFields returns the label-derived fields of a remote agent's container. +func agentLabelFields(labels map[string]string) *Container { + c := &Container{ + OrchestrationGroup: OrchestrationGroupFromLabels(labels), + AlertSeverity: SeverityWarning, + RestartThreshold: 3, + } + applyAgentLabels(c, labels) + return c +} diff --git a/internal/container/service.go b/internal/container/service.go index 986aa3a6..64918a53 100644 --- a/internal/container/service.go +++ b/internal/container/service.go @@ -208,6 +208,9 @@ func (s *Service) handleStateChange(ctx context.Context, evt ContainerEvent, new } return } + if c.IsIgnored { + return + } previousState := c.State if evt.Replayed { @@ -336,7 +339,7 @@ func (s *Service) handleHealthChange(ctx context.Context, evt ContainerEvent) { s.logger.Error("get container for health change", "external_id", shortID(evt.ExternalID), "error", err) return } - if c == nil { + if c == nil || c.IsIgnored { return } @@ -471,7 +474,9 @@ func (s *Service) Reconcile(ctx context.Context, discoverer RuntimeDiscoverer) e dc.ImageVersion, dc.ImageSource, dc.ImageURL, dc.ImageDescription } - if sc.State == dc.State && metadataChanged { + labelsChanged := sc.adoptLabelFields(dc) + + if sc.State == dc.State && (metadataChanged || labelsChanged) { if err := s.store.UpdateContainer(ctx, sc); err != nil { s.logger.Error("reconcile update", "container_id", sc.ID, "error", err) } @@ -496,9 +501,11 @@ func (s *Service) Reconcile(ctx context.Context, discoverer RuntimeDiscoverer) e s.logger.Error("reconcile update", "container_id", sc.ID, "error", err) } - s.emitEvent(event.ContainerStateChanged, map[string]interface{}{ - "id": sc.ID, "state": dc.State, "previous_state": previousState, "timestamp": now, "agent_id": sc.AgentID, - }) + if !sc.IsIgnored { + s.emitEvent(event.ContainerStateChanged, map[string]interface{}{ + "id": sc.ID, "state": dc.State, "previous_state": previousState, "timestamp": now, "agent_id": sc.AgentID, + }) + } } } @@ -532,7 +539,9 @@ func (s *Service) Reconcile(ctx context.Context, discoverer RuntimeDiscoverer) e } } - s.emitEvent(event.ContainerDiscovered, dc) + if !dc.IsIgnored { + s.emitEvent(event.ContainerDiscovered, dc) + } } } diff --git a/internal/docker/client.go b/internal/docker/client.go index 7fb37e13..c0e4f060 100644 --- a/internal/docker/client.go +++ b/internal/docker/client.go @@ -30,6 +30,9 @@ type Client struct { connected bool proxyLabels bool + + serviceMu sync.Mutex + serviceLabels map[string]serviceLabels } // NewClient creates a new Docker client wrapper. @@ -59,7 +62,11 @@ func (c *Client) SetProxyLabels(enabled bool) { c.proxyLabels = enabled } -func (c *Client) containerLabels(labels map[string]string) map[string]string { +// containerLabels returns the labels maintenant reads for a container: its own, +// over those of its Swarm service, expanded from reverse proxy labels when +// enabled. +func (c *Client) containerLabels(ctx context.Context, labels map[string]string) map[string]string { + labels = withServiceLabels(labels, c.swarmServiceLabels(ctx, labels[labelSwarmServiceID])) if !c.proxyLabels { return labels } diff --git a/internal/docker/discovery.go b/internal/docker/discovery.go index b4c04ce7..3fa5437f 100644 --- a/internal/docker/discovery.go +++ b/internal/docker/discovery.go @@ -18,7 +18,6 @@ import ( ) const ( - labelComposeProject = "com.docker.compose.project" labelComposeService = "com.docker.compose.service" labelComposeWorkingDir = "com.docker.compose.project.working_dir" labelComposeOneOff = "com.docker.compose.oneoff" @@ -81,10 +80,11 @@ func (c *Client) DiscoverAll(ctx context.Context) ([]*cmodel.Container, error) { if IsOneOff(dc.Labels) { continue } - result, err := c.inspectAndMap(ctx, dc, now) + labels := c.containerLabels(ctx, dc.Labels) + result, err := c.inspectAndMap(ctx, dc, labels, now) if err != nil { c.logger.Warn("failed to inspect container", "docker_id", dc.ID[:12], "error", err) - containers = append(containers, mapFromList(dc, now)) + containers = append(containers, mapFromList(dc, labels, now)) continue } containers = append(containers, result.Container) @@ -108,18 +108,19 @@ func (c *Client) DiscoverAllWithLabels(ctx context.Context) ([]*DiscoveryResult, if IsOneOff(dc.Labels) { continue } - result, err := c.inspectAndMap(ctx, dc, now) + labels := c.containerLabels(ctx, dc.Labels) + result, err := c.inspectAndMap(ctx, dc, labels, now) if err != nil { c.logger.Warn("failed to inspect container", "docker_id", dc.ID[:12], "error", err) results = append(results, &DiscoveryResult{ - Container: mapFromList(dc, now), - Labels: c.containerLabels(dc.Labels), + Container: mapFromList(dc, labels, now), + Labels: labels, }) continue } results = append(results, &DiscoveryResult{ Container: result.Container, - Labels: c.containerLabels(dc.Labels), + Labels: labels, SecurityConfig: result.SecurityConfig, }) } @@ -159,14 +160,14 @@ type inspectResult struct { } // inspectAndMap calls ContainerInspect and maps the result to our domain model. -func (c *Client) inspectAndMap(ctx context.Context, dc container.Summary, now time.Time) (*inspectResult, error) { +func (c *Client) inspectAndMap(ctx context.Context, dc container.Summary, labels map[string]string, now time.Time) (*inspectResult, error) { res, err := c.cli.ContainerInspect(ctx, dc.ID, client.ContainerInspectOptions{}) if err != nil { return nil, fmt.Errorf("inspect %s: %w", dc.ID[:12], err) } info := res.Container - cm := mapFromList(dc, now) + cm := mapFromList(dc, labels, now) // Health check info from inspect if info.Config != nil && info.Config.Healthcheck != nil && len(info.Config.Healthcheck.Test) > 0 { @@ -218,8 +219,9 @@ func extractSecurityConfig(hc *container.HostConfig) *SecurityConfig { return cfg } -// mapFromList creates a Container from the docker ContainerList response. -func mapFromList(dc container.Summary, now time.Time) *cmodel.Container { +// mapFromList creates a Container from the docker ContainerList response and +// the container's effective labels. +func mapFromList(dc container.Summary, labels map[string]string, now time.Time) *cmodel.Container { name := "" if len(dc.Names) > 0 { name = dc.Names[0] @@ -239,9 +241,9 @@ func mapFromList(dc container.Summary, now time.Time) *cmodel.Container { Name: name, Image: dc.Image, State: state, - OrchestrationGroup: dc.Labels[labelComposeProject], - OrchestrationUnit: dc.Labels[labelComposeService], - ComposeWorkingDir: dc.Labels[labelComposeWorkingDir], + OrchestrationGroup: cmodel.OrchestrationGroupFromLabels(labels), + OrchestrationUnit: labels[labelComposeService], + ComposeWorkingDir: labels[labelComposeWorkingDir], RuntimeType: "docker", PodCount: 1, ReadyCount: readyCount, @@ -251,8 +253,8 @@ func mapFromList(dc container.Summary, now time.Time) *cmodel.Container { LastStateChangeAt: now, } - applyLabels(cm, dc.Labels) - cm.ApplyImageLabels(dc.Labels) + applyLabels(cm, labels) + cm.ApplyImageLabels(labels) return cm } diff --git a/internal/docker/events.go b/internal/docker/events.go index 09a59ae0..fafb09e4 100644 --- a/internal/docker/events.go +++ b/internal/docker/events.go @@ -75,7 +75,7 @@ func (c *Client) StreamEvents(ctx context.Context) <-chan ContainerEvent { continue } if evt.ResourceType == "container" { - evt.Labels = c.containerLabels(evt.Labels) + evt.Labels = c.containerLabels(ctx, evt.Labels) } select { diff --git a/internal/docker/service_labels.go b/internal/docker/service_labels.go new file mode 100644 index 00000000..0721d8d7 --- /dev/null +++ b/internal/docker/service_labels.go @@ -0,0 +1,77 @@ +// Copyright 2026 Benjamin Touchard (Kolapsis) +// SPDX-License-Identifier: Apache-2.0 + +package docker + +import ( + "context" + "time" + + "github.com/moby/moby/client" +) + +const ( + labelSwarmServiceID = "com.docker.swarm.service.id" + + serviceLabelsTTL = 30 * time.Second + serviceLabelsIdle = 10 * time.Minute +) + +type serviceLabels struct { + labels map[string]string + fetchedAt time.Time +} + +// swarmServiceLabels returns the labels of a Swarm service, the deploy.labels +// of a stack, cached for serviceLabelsTTL. A node that cannot read services, a +// worker, gets none; a failed refresh keeps the labels last read. +func (c *Client) swarmServiceLabels(ctx context.Context, serviceID string) map[string]string { + if serviceID == "" { + return nil + } + now := time.Now() + + c.serviceMu.Lock() + cached, ok := c.serviceLabels[serviceID] + c.serviceMu.Unlock() + if ok && now.Sub(cached.fetchedAt) < serviceLabelsTTL { + return cached.labels + } + + labels := cached.labels + res, err := c.cli.ServiceInspect(ctx, serviceID, client.ServiceInspectOptions{}) + if err != nil { + c.logger.Debug("swarm service labels unavailable", "service_id", serviceID, "error", err) + } else { + labels = res.Service.Spec.Labels + } + + c.serviceMu.Lock() + defer c.serviceMu.Unlock() + if c.serviceLabels == nil { + c.serviceLabels = make(map[string]serviceLabels) + } + for id, e := range c.serviceLabels { + if now.Sub(e.fetchedAt) > serviceLabelsIdle { + delete(c.serviceLabels, id) + } + } + c.serviceLabels[serviceID] = serviceLabels{labels: labels, fetchedAt: now} + return labels +} + +// withServiceLabels lays a container's labels over those of its service: on a +// key both set, the container's value wins. +func withServiceLabels(labels, service map[string]string) map[string]string { + if len(service) == 0 { + return labels + } + merged := make(map[string]string, len(service)+len(labels)) + for k, v := range service { + merged[k] = v + } + for k, v := range labels { + merged[k] = v + } + return merged +} diff --git a/internal/docker/service_labels_test.go b/internal/docker/service_labels_test.go new file mode 100644 index 00000000..fb86ef36 --- /dev/null +++ b/internal/docker/service_labels_test.go @@ -0,0 +1,146 @@ +// Copyright 2026 Benjamin Touchard (Kolapsis) +// SPDX-License-Identifier: Apache-2.0 + +package docker + +import ( + "context" + "errors" + "testing" + "time" + + "github.com/moby/moby/api/types/container" + "github.com/moby/moby/api/types/events" + "github.com/moby/moby/api/types/swarm" + "github.com/moby/moby/client" + "github.com/stretchr/testify/assert" + "github.com/stretchr/testify/require" +) + +type swarmFakeAPI struct { + fakeAPI + services map[string]map[string]string + inspectErr error + inspects int +} + +func (f *swarmFakeAPI) ServiceInspect(_ context.Context, id string, _ client.ServiceInspectOptions) (client.ServiceInspectResult, error) { + f.inspects++ + if f.inspectErr != nil { + return client.ServiceInspectResult{}, f.inspectErr + } + return client.ServiceInspectResult{Service: swarm.Service{ + ID: id, + Spec: swarm.ServiceSpec{Annotations: swarm.Annotations{Labels: f.services[id]}}, + }}, nil +} + +func taskSummary() container.Summary { + return container.Summary{ + ID: "0123456789abcdef", + Names: []string{"/prod_web.1.x1y2z3"}, + State: "running", + Labels: map[string]string{ + "com.docker.swarm.service.id": "svc1", + "com.docker.swarm.service.name": "prod_web", + "com.docker.stack.namespace": "prod", + "maintenant.alert.severity": "info", + }, + } +} + +func TestDiscoverAllWithLabels_ReadsSwarmServiceLabels(t *testing.T) { + api := &swarmFakeAPI{ + fakeAPI: fakeAPI{list: []container.Summary{taskSummary()}}, + services: map[string]map[string]string{"svc1": { + "maintenant.alert.severity": "critical", + "maintenant.group": "shop", + "traefik.http.routers.web.rule": "Host(`shop.example.com`)", + "traefik.http.routers.web.tls.certresolver": "le", + "maintenant.update.tag-include": `^\d+\.\d+$`, + "com.docker.stack.namespace": "prod", + "com.docker.swarm.service.name": "service-value", + "maintenant.endpoint.http.expected-status": "200", + "maintenant.tls.certificates": "api.example.com", + "maintenant.alert.restart_threshold": "5", + "traefik.http.routers.web.entrypoints": "websecure", + "traefik.http.services.web.loadbalancer.port": "8080", + }}, + } + c := &Client{cli: api, logger: newFakeClient(&api.fakeAPI).logger} + c.SetProxyLabels(true) + + results, err := c.DiscoverAllWithLabels(context.Background()) + require.NoError(t, err) + require.Len(t, results, 1) + got := results[0] + + assert.Equal(t, "prod", got.Container.OrchestrationGroup, "grouped by stack without a Compose project") + assert.Equal(t, "shop", got.Container.CustomGroup) + assert.Equal(t, 5, got.Container.RestartThreshold) + assert.Equal(t, "info", string(got.Container.AlertSeverity), "the container's own label wins") + assert.Equal(t, "prod_web", got.Labels["com.docker.swarm.service.name"]) + assert.Equal(t, `^\d+\.\d+$`, got.Labels["maintenant.update.tag-include"]) + assert.Equal(t, "https://shop.example.com", got.Labels["maintenant.endpoint.0.http"], + "Traefik labels set in deploy.labels feed the proxy expansion") + assert.Contains(t, got.Labels["maintenant.tls.certificates"], "api.example.com") + + _, err = c.DiscoverAll(context.Background()) + require.NoError(t, err) + assert.Equal(t, 1, api.inspects, "service labels are cached between passes") +} + +func TestDiscoverAll_IgnoreInDeployLabels(t *testing.T) { + api := &swarmFakeAPI{ + fakeAPI: fakeAPI{list: []container.Summary{taskSummary()}}, + services: map[string]map[string]string{"svc1": {"maintenant.ignore": "true"}}, + } + c := &Client{cli: api, logger: newFakeClient(&api.fakeAPI).logger} + + containers, err := c.DiscoverAll(context.Background()) + require.NoError(t, err) + require.Len(t, containers, 1) + assert.True(t, containers[0].IsIgnored) +} + +func TestDiscoverAll_WorkerWithoutServiceAccess(t *testing.T) { + api := &swarmFakeAPI{ + fakeAPI: fakeAPI{list: []container.Summary{taskSummary()}}, + inspectErr: errors.New("This node is not a swarm manager"), + } + c := &Client{cli: api, logger: newFakeClient(&api.fakeAPI).logger} + + containers, err := c.DiscoverAll(context.Background()) + require.NoError(t, err) + require.Len(t, containers, 1) + assert.False(t, containers[0].IsIgnored) + assert.Equal(t, "prod", containers[0].OrchestrationGroup) + assert.Equal(t, "info", string(containers[0].AlertSeverity)) +} + +func TestStreamEvents_ReadsSwarmServiceLabels(t *testing.T) { + api := &swarmFakeAPI{ + fakeAPI: fakeAPI{events: make(chan events.Message, 1)}, + services: map[string]map[string]string{"svc1": {"maintenant.ignore": "true"}}, + } + api.events <- events.Message{ + Type: events.ContainerEventType, + Action: "start", + Actor: events.Actor{ + ID: "ctr1", + Attributes: map[string]string{"name": "prod_web.1.x1y2z3", "com.docker.swarm.service.id": "svc1"}, + }, + Time: 1, + } + c := &Client{cli: api, logger: newFakeClient(&api.fakeAPI).logger} + + ctx, cancel := context.WithCancel(context.Background()) + defer cancel() + select { + case evt := <-c.StreamEvents(ctx): + assert.Equal(t, "true", evt.Labels["maintenant.ignore"]) + assert.Equal(t, "prod_web.1.x1y2z3", evt.Labels["name"]) + case <-time.After(2 * time.Second): + t.Fatal("no event received") + } +} diff --git a/internal/endpoint/labels.go b/internal/endpoint/labels.go index 66db43b5..02611d38 100644 --- a/internal/endpoint/labels.go +++ b/internal/endpoint/labels.go @@ -12,6 +12,8 @@ import ( "strconv" "strings" "time" + + "github.com/kolapsis/maintenant/internal/container" ) const ( @@ -45,8 +47,12 @@ func (e *LabelParseError) Error() string { } // ParseEndpointLabels extracts endpoint definitions from a Docker container's labels. -// Returns parsed endpoints and any configuration errors encountered. +// Returns parsed endpoints and any configuration errors encountered; an ignored +// container declares none. func ParseEndpointLabels(labels map[string]string, logger *slog.Logger) ([]*ParsedEndpoint, []*LabelParseError) { + if container.IgnoredByLabels(labels) { + return nil, nil + } endpointMap := make(map[int]*ParsedEndpoint) globalConfig := make(map[string]string) indexedConfigs := make(map[int]map[string]string) diff --git a/internal/endpoint/service_test.go b/internal/endpoint/service_test.go index 6678e9ad..43a808e2 100644 --- a/internal/endpoint/service_test.go +++ b/internal/endpoint/service_test.go @@ -1089,3 +1089,39 @@ func TestService_ProcessCheckResult_DuplicateReplayDoesNotInflateCounters(t *tes assert.Equal(t, first.ConsecutiveFailures, second.ConsecutiveFailures, "the same probe delivered twice must count once") } + +func TestParseEndpointLabels_IgnoredContainerDeclaresNone(t *testing.T) { + parsed, errs := ParseEndpointLabels(map[string]string{ + "maintenant.ignore": "true", + "maintenant.endpoint.http": "http://web:8080/health", + "maintenant.endpoint.tcp": "not a target", + }, noopLogger()) + assert.Empty(t, parsed) + assert.Empty(t, errs) +} + +func TestService_SyncEndpoints_IgnoredContainerLosesItsEndpoints(t *testing.T) { + store := newMemStore() + svc := newService(store) + ctx := context.Background() + labels := map[string]string{"maintenant.endpoint.http": "http://web:8080/health"} + + svc.SyncEndpoints(ctx, "web", "container-1", labels, "", "") + eps, err := store.ListEndpointsByExternalID(ctx, uid.LocalAgent, "container-1") + require.NoError(t, err) + require.Len(t, eps, 1) + + labels["maintenant.ignore"] = "true" + svc.SyncEndpoints(ctx, "web", "container-1", labels, "", "") + ep, err := store.GetEndpointByID(ctx, eps[0].ID) + require.NoError(t, err) + assert.False(t, ep.Active, "an ignored container keeps no endpoint") + + svc.SyncAgentEndpoints(ctx, "agent-1", "api", "container-2", map[string]string{ + "maintenant.ignore": "1", + "maintenant.endpoint.http": "http://api:8080/health", + }) + agentEps, err := store.ListEndpointsByExternalID(ctx, "agent-1", "container-2") + require.NoError(t, err) + assert.Empty(t, agentEps) +} diff --git a/internal/kubernetes/services.go b/internal/kubernetes/services.go new file mode 100644 index 00000000..caf371a9 --- /dev/null +++ b/internal/kubernetes/services.go @@ -0,0 +1,166 @@ +// Copyright 2026 Benjamin Touchard (Kolapsis) +// SPDX-License-Identifier: Apache-2.0 + +package kubernetes + +import ( + "context" + "fmt" + "strings" + + corev1 "k8s.io/api/core/v1" + k8serrors "k8s.io/apimachinery/pkg/api/errors" + metav1 "k8s.io/apimachinery/pkg/apis/meta/v1" + "k8s.io/apimachinery/pkg/labels" + "k8s.io/apimachinery/pkg/util/intstr" +) + +// ServiceExposure is one port a LoadBalancer or NodePort Service opens outside +// the cluster, attributed to a workload its selector matches. +type ServiceExposure struct { + WorkloadID string // the workload's external id, as discoverAll mints it + Service string // namespace/name + ServiceType string + Port int + TargetPort int // 0 when the Service targets a named port + NodePort int + Protocol string +} + +type selectableWorkload struct { + id string + namespace string + labels map[string]string +} + +// ListServiceExposures returns the ports LoadBalancer and NodePort Services +// expose, one entry per workload each Service selects. +func (r *Runtime) ListServiceExposures(ctx context.Context) ([]ServiceExposure, error) { + list, err := r.clientset.CoreV1().Services("").List(ctx, metav1.ListOptions{}) + if err != nil { + return nil, fmt.Errorf("list services: %w", err) + } + + var exposing []corev1.Service + for _, svc := range list.Items { + if !r.nsFilter.IsAllowed(svc.Namespace) || len(svc.Spec.Selector) == 0 { + continue + } + if svc.Spec.Type == corev1.ServiceTypeLoadBalancer || svc.Spec.Type == corev1.ServiceTypeNodePort { + exposing = append(exposing, svc) + } + } + if len(exposing) == 0 { + return nil, nil + } + + workloads, err := r.selectableWorkloads(ctx) + if err != nil { + return nil, err + } + + var out []ServiceExposure + for _, svc := range exposing { + selector := labels.SelectorFromSet(svc.Spec.Selector) + for _, w := range workloads { + if w.namespace != svc.Namespace || !selector.Matches(labels.Set(w.labels)) { + continue + } + for _, p := range svc.Spec.Ports { + out = append(out, ServiceExposure{ + WorkloadID: w.id, + Service: svc.Namespace + "/" + svc.Name, + ServiceType: string(svc.Spec.Type), + Port: int(p.Port), + TargetPort: numericTargetPort(p.TargetPort), + NodePort: int(p.NodePort), + Protocol: servicePortProtocol(p.Protocol), + }) + } + } + } + return out, nil +} + +// selectableWorkloads lists what a Service selector can match: the pod +// templates of the controllers and the labels of bare pods. A kind the RBAC +// denies is left out, as in discoverAll. +func (r *Runtime) selectableWorkloads(ctx context.Context) ([]selectableWorkload, error) { + var out []selectableWorkload + keep := func(ns string) bool { return r.nsFilter.IsAllowed(ns) } + skip := func(kind string, err error) error { + if k8serrors.IsForbidden(err) { + r.logger.Warn("RBAC: forbidden to list "+kind+", its exposure is not analysed", "error", err) + return nil + } + return fmt.Errorf("list %s: %w", kind, err) + } + + deployments, err := r.clientset.AppsV1().Deployments("").List(ctx, metav1.ListOptions{}) + if err != nil { + if err := skip("deployments", err); err != nil { + return nil, err + } + } else { + for _, d := range deployments.Items { + if keep(d.Namespace) { + out = append(out, selectableWorkload{fmt.Sprintf("%s/Deployment/%s", d.Namespace, d.Name), d.Namespace, d.Spec.Template.Labels}) + } + } + } + + statefulSets, err := r.clientset.AppsV1().StatefulSets("").List(ctx, metav1.ListOptions{}) + if err != nil { + if err := skip("statefulsets", err); err != nil { + return nil, err + } + } else { + for _, s := range statefulSets.Items { + if keep(s.Namespace) { + out = append(out, selectableWorkload{fmt.Sprintf("%s/StatefulSet/%s", s.Namespace, s.Name), s.Namespace, s.Spec.Template.Labels}) + } + } + } + + daemonSets, err := r.clientset.AppsV1().DaemonSets("").List(ctx, metav1.ListOptions{}) + if err != nil { + if err := skip("daemonsets", err); err != nil { + return nil, err + } + } else { + for _, d := range daemonSets.Items { + if keep(d.Namespace) { + out = append(out, selectableWorkload{fmt.Sprintf("%s/DaemonSet/%s", d.Namespace, d.Name), d.Namespace, d.Spec.Template.Labels}) + } + } + } + + pods, err := r.clientset.CoreV1().Pods("").List(ctx, metav1.ListOptions{}) + if err != nil { + if err := skip("pods", err); err != nil { + return nil, err + } + } else { + for i := range pods.Items { + p := &pods.Items[i] + if keep(p.Namespace) && !hasControllerOwner(p) { + out = append(out, selectableWorkload{fmt.Sprintf("%s/%s", p.Namespace, p.Name), p.Namespace, p.Labels}) + } + } + } + return out, nil +} + +func numericTargetPort(p intstr.IntOrString) int { + if p.Type == intstr.Int { + return int(p.IntVal) + } + return 0 +} + +func servicePortProtocol(p corev1.Protocol) string { + if p == "" { + return "tcp" + } + return strings.ToLower(string(p)) +} diff --git a/internal/kubernetes/services_test.go b/internal/kubernetes/services_test.go new file mode 100644 index 00000000..6fbacf43 --- /dev/null +++ b/internal/kubernetes/services_test.go @@ -0,0 +1,107 @@ +// Copyright 2026 Benjamin Touchard (Kolapsis) +// SPDX-License-Identifier: Apache-2.0 + +package kubernetes + +import ( + "context" + "log/slog" + "testing" + + "github.com/stretchr/testify/assert" + "github.com/stretchr/testify/require" + appsv1 "k8s.io/api/apps/v1" + corev1 "k8s.io/api/core/v1" + metav1 "k8s.io/apimachinery/pkg/apis/meta/v1" + k8sruntime "k8s.io/apimachinery/pkg/runtime" + "k8s.io/apimachinery/pkg/util/intstr" + "k8s.io/client-go/kubernetes/fake" +) + +func serviceRuntime(nsFilter *NamespaceFilter, objects ...k8sruntime.Object) *Runtime { + return &Runtime{ + logger: slog.Default(), + nsFilter: nsFilter, + clientset: fake.NewClientset(objects...), + prevCPU: make(map[string]*cpuPrev), + stopCh: make(chan struct{}), + } +} + +func service(ns, name string, typ corev1.ServiceType, selector map[string]string, ports ...corev1.ServicePort) *corev1.Service { + return &corev1.Service{ + ObjectMeta: metav1.ObjectMeta{Namespace: ns, Name: name}, + Spec: corev1.ServiceSpec{Type: typ, Selector: selector, Ports: ports}, + } +} + +func TestListServiceExposures(t *testing.T) { + web := &appsv1.Deployment{ + ObjectMeta: metav1.ObjectMeta{Namespace: "shop", Name: "web"}, + Spec: appsv1.DeploymentSpec{Template: corev1.PodTemplateSpec{ + ObjectMeta: metav1.ObjectMeta{Labels: map[string]string{"app": "web", "tier": "front"}}, + }}, + } + db := &appsv1.StatefulSet{ + ObjectMeta: metav1.ObjectMeta{Namespace: "shop", Name: "db"}, + Spec: appsv1.StatefulSetSpec{Template: corev1.PodTemplateSpec{ + ObjectMeta: metav1.ObjectMeta{Labels: map[string]string{"app": "db"}}, + }}, + } + debug := &corev1.Pod{ObjectMeta: metav1.ObjectMeta{Namespace: "shop", Name: "debug", Labels: map[string]string{"app": "debug"}}} + otherNS := &appsv1.Deployment{ + ObjectMeta: metav1.ObjectMeta{Namespace: "staging", Name: "web"}, + Spec: appsv1.DeploymentSpec{Template: corev1.PodTemplateSpec{ + ObjectMeta: metav1.ObjectMeta{Labels: map[string]string{"app": "web"}}, + }}, + } + + rt := serviceRuntime(NewNamespaceFilter("", ""), + web, db, debug, otherNS, + service("shop", "web-lb", corev1.ServiceTypeLoadBalancer, map[string]string{"app": "web"}, + corev1.ServicePort{Port: 443, TargetPort: intstr.FromInt32(8443), Protocol: corev1.ProtocolTCP}), + service("shop", "db-np", corev1.ServiceTypeNodePort, map[string]string{"app": "db"}, + corev1.ServicePort{Port: 5432, TargetPort: intstr.FromString("pg"), NodePort: 31432}), + service("shop", "debug-np", corev1.ServiceTypeNodePort, map[string]string{"app": "debug"}, + corev1.ServicePort{Port: 8080, NodePort: 30080, Protocol: corev1.ProtocolUDP}), + service("shop", "web-internal", corev1.ServiceTypeClusterIP, map[string]string{"app": "web"}, + corev1.ServicePort{Port: 80}), + service("shop", "no-selector", corev1.ServiceTypeLoadBalancer, nil, + corev1.ServicePort{Port: 80}), + ) + + exposures, err := rt.ListServiceExposures(context.Background()) + require.NoError(t, err) + + byWorkload := map[string]ServiceExposure{} + for _, e := range exposures { + byWorkload[e.WorkloadID] = e + } + require.Len(t, byWorkload, 3, "the ClusterIP service, the selector-less one and the other namespace expose nothing") + + assert.Equal(t, ServiceExposure{ + WorkloadID: "shop/Deployment/web", Service: "shop/web-lb", ServiceType: "LoadBalancer", + Port: 443, TargetPort: 8443, Protocol: "tcp", + }, byWorkload["shop/Deployment/web"]) + assert.Equal(t, ServiceExposure{ + WorkloadID: "shop/StatefulSet/db", Service: "shop/db-np", ServiceType: "NodePort", + Port: 5432, NodePort: 31432, Protocol: "tcp", + }, byWorkload["shop/StatefulSet/db"]) + assert.Equal(t, "udp", byWorkload["shop/debug"].Protocol) +} + +func TestListServiceExposures_NamespaceFilter(t *testing.T) { + web := &appsv1.Deployment{ + ObjectMeta: metav1.ObjectMeta{Namespace: "shop", Name: "web"}, + Spec: appsv1.DeploymentSpec{Template: corev1.PodTemplateSpec{ + ObjectMeta: metav1.ObjectMeta{Labels: map[string]string{"app": "web"}}, + }}, + } + rt := serviceRuntime(NewNamespaceFilter("", "shop"), web, + service("shop", "web-lb", corev1.ServiceTypeLoadBalancer, map[string]string{"app": "web"}, + corev1.ServicePort{Port: 443})) + + exposures, err := rt.ListServiceExposures(context.Background()) + require.NoError(t, err) + assert.Empty(t, exposures) +} diff --git a/internal/resource/agent_event.go b/internal/resource/agent_event.go index 9e94d4a8..0ff4e5f2 100644 --- a/internal/resource/agent_event.go +++ b/internal/resource/agent_event.go @@ -32,7 +32,7 @@ func (s *Service) HandleAgentEvent(ctx context.Context, agentID string, ev *agen } c, err := s.containerSvc.GetContainerByExternalID(ctx, agentID, containerExternalID) - if err != nil || c == nil { + if err != nil || c == nil || c.IsIgnored { return err } diff --git a/internal/resource/agent_event_test.go b/internal/resource/agent_event_test.go index 00ae913e..2eb15b86 100644 --- a/internal/resource/agent_event_test.go +++ b/internal/resource/agent_event_test.go @@ -74,6 +74,21 @@ func TestHandleAgentEvent_SkipsWhenContainerUnknown(t *testing.T) { assert.Empty(t, rstore.snapshots, "no snapshot when container not yet known") } +func TestHandleAgentEvent_SkipsIgnoredContainer(t *testing.T) { + extID := "abc123def4567890" + c := &container.Container{ + ID: uid.Container(uid.Agent("agent-9"), extID), + ExternalID: extID, AgentID: "agent-9", Name: "demo", IsIgnored: true, + } + rstore := newMockResourceStore() + svc := newTestService(rstore, buildContainerSvc(newMockContainerStore(c)), nil) + + require.NoError(t, svc.HandleAgentEvent(context.Background(), "agent-9", &agentpb.ResourceSample{ + ContainerId: extID, CpuPercent: 3, + }, agentevent.Meta{ObservedAt: time.Now()})) + assert.Empty(t, rstore.snapshots, "an ignored container is left out of resource collection") +} + func TestHandleAgentEvent_UsesRowIDNotDerivedID(t *testing.T) { extID := "abc123def4567890" // An inherited row: its primary key does not derive from its current agent. diff --git a/internal/security/analyzer.go b/internal/security/analyzer.go index 4628f3a3..59c50f90 100644 --- a/internal/security/analyzer.go +++ b/internal/security/analyzer.go @@ -101,6 +101,67 @@ func analyzePortBindings(containerID string, containerName string, bindings []Po return insights } +// Kubernetes Service types that open a port outside the cluster. +const ( + ServiceTypeLoadBalancer = "LoadBalancer" + ServiceTypeNodePort = "NodePort" +) + +// ServicePort is one port a LoadBalancer or NodePort Service exposes for a workload. +type ServicePort struct { + Service string // namespace/name + Type string // ServiceTypeLoadBalancer or ServiceTypeNodePort + Port int + TargetPort int // 0 when the Service targets a named port + NodePort int + Protocol string +} + +// AnalyzeKubernetes returns the insights of the ports Services expose for a workload. +func AnalyzeKubernetes(containerID string, containerName string, ports []ServicePort, now time.Time) []Insight { + var insights []Insight + for _, p := range ports { + port := p.TargetPort + if port == 0 { + port = p.Port + } + details := map[string]any{"port": port, "protocol": p.Protocol, "service": p.Service, "service_type": p.Type} + if p.NodePort != 0 { + details["node_port"] = p.NodePort + } + insight := Insight{ + Severity: SeverityCritical, + ContainerID: containerID, + ContainerName: containerName, + Details: details, + DetectedAt: now, + } + + switch dbType, isDB := knownDatabasePorts[port]; { + case isDB: + insight.Type = DatabasePortExposed + insight.Title = "Database port publicly exposed" + insight.Description = fmt.Sprintf("%s port %d is reachable from outside the cluster through %s service %s.", + dbType, port, p.Type, p.Service) + details["database_type"] = dbType + case p.Type == ServiceTypeLoadBalancer: + insight.Type = ServiceLoadBalancer + insight.Title = "Port exposed by a LoadBalancer service" + insight.Description = fmt.Sprintf("Port %d/%s is reachable from outside the cluster through LoadBalancer service %s.", + port, p.Protocol, p.Service) + case p.Type == ServiceTypeNodePort: + insight.Type = ServiceNodePort + insight.Title = "Port exposed on every node" + insight.Description = fmt.Sprintf("Port %d/%s is published on port %d of every node by NodePort service %s.", + port, p.Protocol, p.NodePort, p.Service) + default: + continue + } + insights = append(insights, insight) + } + return insights +} + func isExposedOnAllInterfaces(hostIP string) bool { return hostIP == "" || hostIP == "0.0.0.0" || hostIP == "::" } diff --git a/internal/security/analyzer_kubernetes_test.go b/internal/security/analyzer_kubernetes_test.go new file mode 100644 index 00000000..3236489d --- /dev/null +++ b/internal/security/analyzer_kubernetes_test.go @@ -0,0 +1,56 @@ +// Copyright 2026 Benjamin Touchard (Kolapsis) +// SPDX-License-Identifier: Apache-2.0 + +package security + +import ( + "testing" + + "github.com/stretchr/testify/assert" + "github.com/stretchr/testify/require" +) + +func TestAnalyzeKubernetes_LoadBalancer(t *testing.T) { + insights := AnalyzeKubernetes("w1", "web", []ServicePort{ + {Service: "shop/web", Type: ServiceTypeLoadBalancer, Port: 443, TargetPort: 8443, Protocol: "tcp"}, + }, testNow) + + require.Len(t, insights, 1) + assert.Equal(t, ServiceLoadBalancer, insights[0].Type) + assert.Equal(t, SeverityCritical, insights[0].Severity, "as critical as a Docker port bound on every interface") + assert.Equal(t, 8443, insights[0].Details["port"]) + assert.Equal(t, "tcp", insights[0].Details["protocol"]) + assert.Equal(t, "shop/web", insights[0].Details["service"]) + assert.Equal(t, "8443/tcp", InsightFindingKey(insights[0])) +} + +func TestAnalyzeKubernetes_NodePort(t *testing.T) { + insights := AnalyzeKubernetes("w1", "web", []ServicePort{ + {Service: "shop/web", Type: ServiceTypeNodePort, Port: 80, NodePort: 30080, Protocol: "tcp"}, + }, testNow) + + require.Len(t, insights, 1) + assert.Equal(t, ServiceNodePort, insights[0].Type) + assert.Equal(t, SeverityCritical, insights[0].Severity) + assert.Equal(t, 80, insights[0].Details["port"], "a named target port falls back on the service port") + assert.Equal(t, 30080, insights[0].Details["node_port"]) +} + +func TestAnalyzeKubernetes_DatabasePort(t *testing.T) { + insights := AnalyzeKubernetes("w1", "db", []ServicePort{ + {Service: "data/pg", Type: ServiceTypeNodePort, Port: 5432, TargetPort: 5432, NodePort: 31432, Protocol: "tcp"}, + {Service: "data/cache", Type: ServiceTypeLoadBalancer, Port: 6380, TargetPort: 6379, Protocol: "tcp"}, + }, testNow) + + require.Len(t, insights, 2) + for _, i := range insights { + assert.Equal(t, DatabasePortExposed, i.Type) + assert.Equal(t, SeverityCritical, i.Severity) + } + assert.Equal(t, "PostgreSQL", insights[0].Details["database_type"]) + assert.Equal(t, "Redis", insights[1].Details["database_type"]) +} + +func TestAnalyzeKubernetes_NoServiceNoInsight(t *testing.T) { + assert.Empty(t, AnalyzeKubernetes("w1", "web", nil, testNow)) +} diff --git a/internal/security/model.go b/internal/security/model.go index 93beda0a..93021191 100644 --- a/internal/security/model.go +++ b/internal/security/model.go @@ -15,7 +15,6 @@ const ( HostNetworkMode InsightType = "host_network_mode" ServiceLoadBalancer InsightType = "service_load_balancer" ServiceNodePort InsightType = "service_node_port" - MissingNetworkPolicy InsightType = "missing_network_policy" ) // Severity levels for security insights. diff --git a/internal/store/containers.go b/internal/store/containers.go index 49b9fbaa..3e248af7 100644 --- a/internal/store/containers.go +++ b/internal/store/containers.go @@ -333,8 +333,14 @@ func (s *ContainerStore) DeleteTransitionsBefore(ctx context.Context, before tim return deleted, err } +// deleteTransitionsBefore keeps the latest transition of each container, +// however old: it is the state the container is still in, and uptime starts +// from it. func (s *ContainerStore) deleteTransitionsBefore(ctx context.Context, before time.Time, o batchOpts) (int64, bool, error) { - return deleteRowsBefore(ctx, s.writer, o, "state_transitions", "timestamp", before) + return deleteRowsWhere(ctx, s.writer, o, "state_transitions", + `timestamp < ? AND EXISTS (SELECT 1 FROM state_transitions later + WHERE later.container_id = state_transitions.container_id AND later.timestamp > state_transitions.timestamp)`, + before.Unix()) } func (s *ContainerStore) DeleteArchivedContainersBefore(ctx context.Context, before time.Time) (int64, error) { diff --git a/internal/store/copy.go b/internal/store/copy.go index fe1ebd11..0aacbda1 100644 --- a/internal/store/copy.go +++ b/internal/store/copy.go @@ -115,7 +115,8 @@ var leftBehindGroups = []leftBehindGroup{ "kubernetes_nodes", "kubernetes_events"}, "re-sent whole, the full inventory passes every 30s"}, {"Check history", []string{"check_results", "cert_check_results", "cert_chain_entries", - "heartbeat_pings", "heartbeat_executions"}, "starts again, fills itself"}, + "heartbeat_pings", "heartbeat_executions", + "endpoint_uptime_daily", "heartbeat_uptime_daily", "container_uptime_daily"}, "starts again, fills itself"}, {"Resource history", []string{"resource_snapshots", "resource_hourly", "resource_daily"}, "same, and it is the bulk of the volume"}, {"State history", []string{"state_transitions"}, "same"}, diff --git a/internal/store/migrations/postgres/35_uptime_daily.down.sql b/internal/store/migrations/postgres/35_uptime_daily.down.sql new file mode 100644 index 00000000..abdaa577 --- /dev/null +++ b/internal/store/migrations/postgres/35_uptime_daily.down.sql @@ -0,0 +1,3 @@ +DROP TABLE IF EXISTS container_uptime_daily; +DROP TABLE IF EXISTS heartbeat_uptime_daily; +DROP TABLE IF EXISTS endpoint_uptime_daily; diff --git a/internal/store/migrations/postgres/35_uptime_daily.up.sql b/internal/store/migrations/postgres/35_uptime_daily.up.sql new file mode 100644 index 00000000..05f9f330 --- /dev/null +++ b/internal/store/migrations/postgres/35_uptime_daily.up.sql @@ -0,0 +1,26 @@ +CREATE TABLE endpoint_uptime_daily ( + id TEXT PRIMARY KEY NOT NULL, + endpoint_id TEXT NOT NULL REFERENCES endpoints(id) ON DELETE CASCADE, + day BIGINT NOT NULL, + uptime_percent DOUBLE PRECISION NOT NULL, + incident_count INTEGER NOT NULL, + UNIQUE(endpoint_id, day) +); + +CREATE TABLE heartbeat_uptime_daily ( + id TEXT PRIMARY KEY NOT NULL, + heartbeat_id TEXT NOT NULL REFERENCES heartbeats(id) ON DELETE CASCADE, + day BIGINT NOT NULL, + uptime_percent DOUBLE PRECISION NOT NULL, + incident_count INTEGER NOT NULL, + UNIQUE(heartbeat_id, day) +); + +CREATE TABLE container_uptime_daily ( + id TEXT PRIMARY KEY NOT NULL, + container_id TEXT NOT NULL REFERENCES containers(id) ON DELETE CASCADE, + day BIGINT NOT NULL, + uptime_percent DOUBLE PRECISION NOT NULL, + incident_count INTEGER NOT NULL, + UNIQUE(container_id, day) +); diff --git a/internal/store/migrations/sqlite/35_uptime_daily.down.sql b/internal/store/migrations/sqlite/35_uptime_daily.down.sql new file mode 100644 index 00000000..abdaa577 --- /dev/null +++ b/internal/store/migrations/sqlite/35_uptime_daily.down.sql @@ -0,0 +1,3 @@ +DROP TABLE IF EXISTS container_uptime_daily; +DROP TABLE IF EXISTS heartbeat_uptime_daily; +DROP TABLE IF EXISTS endpoint_uptime_daily; diff --git a/internal/store/migrations/sqlite/35_uptime_daily.up.sql b/internal/store/migrations/sqlite/35_uptime_daily.up.sql new file mode 100644 index 00000000..369cd98e --- /dev/null +++ b/internal/store/migrations/sqlite/35_uptime_daily.up.sql @@ -0,0 +1,26 @@ +CREATE TABLE endpoint_uptime_daily ( + id TEXT PRIMARY KEY NOT NULL, + endpoint_id TEXT NOT NULL REFERENCES endpoints(id) ON DELETE CASCADE, + day BIGINT NOT NULL, + uptime_percent REAL NOT NULL, + incident_count INTEGER NOT NULL, + UNIQUE(endpoint_id, day) +); + +CREATE TABLE heartbeat_uptime_daily ( + id TEXT PRIMARY KEY NOT NULL, + heartbeat_id TEXT NOT NULL REFERENCES heartbeats(id) ON DELETE CASCADE, + day BIGINT NOT NULL, + uptime_percent REAL NOT NULL, + incident_count INTEGER NOT NULL, + UNIQUE(heartbeat_id, day) +); + +CREATE TABLE container_uptime_daily ( + id TEXT PRIMARY KEY NOT NULL, + container_id TEXT NOT NULL REFERENCES containers(id) ON DELETE CASCADE, + day BIGINT NOT NULL, + uptime_percent REAL NOT NULL, + incident_count INTEGER NOT NULL, + UNIQUE(container_id, day) +); diff --git a/internal/store/migrations_postgres_test.go b/internal/store/migrations_postgres_test.go index 7abf8d08..b708062c 100644 --- a/internal/store/migrations_postgres_test.go +++ b/internal/store/migrations_postgres_test.go @@ -111,10 +111,11 @@ func TestMigratePostgres_ConcurrentCatchUp(t *testing.T) { "DROP TABLE cve_evaluations", // 31 "ALTER TABLE containers DROP COLUMN image_version, DROP COLUMN image_source, DROP COLUMN image_url, DROP COLUMN image_description", // 32 "ALTER TABLE agents DROP COLUMN os_id, DROP COLUMN os_version_id, DROP COLUMN os_pretty_name, DROP COLUMN os_source, DROP COLUMN os_unavailable_reason, DROP COLUMN os_reported_at", // 33 - "DROP TABLE outbound_heartbeats", // 34 - "ALTER TABLE containers ADD COLUMN alert_channels TEXT", // 36 - "ALTER TABLE escalation_policies ADD COLUMN tags_json TEXT NOT NULL DEFAULT '[]'", // 36 - "ALTER TABLE alert_triggers ADD COLUMN filter_tags TEXT NOT NULL DEFAULT ''", // 36 + "DROP TABLE outbound_heartbeats", // 34 + "DROP TABLE container_uptime_daily, heartbeat_uptime_daily, endpoint_uptime_daily", // 35 + "ALTER TABLE containers ADD COLUMN alert_channels TEXT", // 36 + "ALTER TABLE escalation_policies ADD COLUMN tags_json TEXT NOT NULL DEFAULT '[]'", // 36 + "ALTER TABLE alert_triggers ADD COLUMN filter_tags TEXT NOT NULL DEFAULT ''", // 36 } { _, err = db.ReadDB().Exec(undo) require.NoError(t, err, undo) diff --git a/internal/store/retention.go b/internal/store/retention.go index a9fb095c..89d80576 100644 --- a/internal/store/retention.go +++ b/internal/store/retention.go @@ -28,6 +28,7 @@ const ( resourceSnapshotRetention = resource.DefaultSnapshotRetention resourceHourlyRetention = 90 * 24 * time.Hour // 90 days resourceDailyRetention = 365 * 24 * time.Hour // 1 year + uptimeDailyRetention = 365 * 24 * time.Hour // the longest range the uptime API serves // Below this many deleted rows an autocheckpoint keeps up on its own and // forcing a truncating checkpoint is just noise. @@ -54,6 +55,7 @@ type RetentionConfig struct { HeartbeatPings time.Duration HeartbeatExecs time.Duration CertCheckResults time.Duration + UptimeDaily time.Duration } // withDefaults fills unset fields and rejects values that would break the loop @@ -92,6 +94,7 @@ func (c RetentionConfig) withDefaults(logger *slog.Logger) RetentionConfig { c.HeartbeatPings = orDefault(c.HeartbeatPings, heartbeatPingRetention) c.HeartbeatExecs = orDefault(c.HeartbeatExecs, heartbeatExecRetention) c.CertCheckResults = orDefault(c.CertCheckResults, certCheckResultRetention) + c.UptimeDaily = orDefault(c.UptimeDaily, uptimeDailyRetention) return c } @@ -134,6 +137,7 @@ type RetentionOpts struct { HeartbeatStore *HeartbeatStore CertificateStore *CertificateStore ResourceStore *ResourceStore + UptimeStore *UptimeDailyStore Config RetentionConfig } @@ -199,12 +203,15 @@ func runRetentionPass(ctx context.Context, store *ContainerStore, db *DB, logger started := time.Now() var pass retentionPass - runCleanup(ctx, store, logger, cfg, &pass) + runCleanup(ctx, store, opts.UptimeStore, logger, cfg, &pass) if opts.EndpointStore != nil { - runEndpointCleanup(ctx, opts.EndpointStore, logger, cfg, &pass) + runEndpointCleanup(ctx, opts.EndpointStore, opts.UptimeStore, logger, cfg, &pass) } if opts.HeartbeatStore != nil { - runHeartbeatCleanup(ctx, opts.HeartbeatStore, logger, cfg, &pass) + runHeartbeatCleanup(ctx, opts.HeartbeatStore, opts.UptimeStore, logger, cfg, &pass) + } + if opts.UptimeStore != nil { + runUptimeDailyCleanup(ctx, opts.UptimeStore, logger, cfg, &pass) } if opts.CertificateStore != nil { runCertificateCleanup(ctx, opts.CertificateStore, logger, cfg, &pass) @@ -227,15 +234,17 @@ func runRetentionPass(ctx context.Context, store *ContainerStore, db *DB, logger return pass.truncated } -func runCleanup(ctx context.Context, store *ContainerStore, logger *slog.Logger, cfg RetentionConfig, pass *retentionPass) { +func runCleanup(ctx context.Context, store *ContainerStore, uptime *UptimeDailyStore, logger *slog.Logger, cfg RetentionConfig, pass *retentionPass) { // Clean old transitions - cutoff := time.Now().Add(-cfg.Transitions) - deleted, truncated, err := store.deleteTransitionsBefore(ctx, cutoff, cfg.batchOpts()) - pass.add(deleted, truncated) - if err != nil { - logger.Error("retention cleanup: transitions", "error", err) - } else if deleted > 0 { - logger.Info("retention cleanup: deleted transitions", "count", deleted) + if rolledUp(ctx, logger, uptime, "containers", cfg.Transitions, (*UptimeDailyStore).rollupContainers) { + cutoff := time.Now().Add(-cfg.Transitions) + deleted, truncated, err := store.deleteTransitionsBefore(ctx, cutoff, cfg.batchOpts()) + pass.add(deleted, truncated) + if err != nil { + logger.Error("retention cleanup: transitions", "error", err) + } else if deleted > 0 { + logger.Info("retention cleanup: deleted transitions", "count", deleted) + } } // Clean old archived containers @@ -249,15 +258,17 @@ func runCleanup(ctx context.Context, store *ContainerStore, logger *slog.Logger, } } -func runHeartbeatCleanup(ctx context.Context, store *HeartbeatStore, logger *slog.Logger, cfg RetentionConfig, pass *retentionPass) { +func runHeartbeatCleanup(ctx context.Context, store *HeartbeatStore, uptime *UptimeDailyStore, logger *slog.Logger, cfg RetentionConfig, pass *retentionPass) { // Clean old heartbeat pings - pingCutoff := time.Now().Add(-cfg.HeartbeatPings) - deleted, truncated, err := store.deletePingsBefore(ctx, pingCutoff, cfg.batchOpts()) - pass.add(deleted, truncated) - if err != nil { - logger.Error("retention cleanup: heartbeat pings", "error", err) - } else if deleted > 0 { - logger.Info("retention cleanup: deleted heartbeat pings", "count", deleted) + if rolledUp(ctx, logger, uptime, "heartbeats", cfg.HeartbeatPings, (*UptimeDailyStore).rollupHeartbeats) { + pingCutoff := time.Now().Add(-cfg.HeartbeatPings) + deleted, truncated, err := store.deletePingsBefore(ctx, pingCutoff, cfg.batchOpts()) + pass.add(deleted, truncated) + if err != nil { + logger.Error("retention cleanup: heartbeat pings", "error", err) + } else if deleted > 0 { + logger.Info("retention cleanup: deleted heartbeat pings", "count", deleted) + } } // Clean old heartbeat executions @@ -311,15 +322,17 @@ func runResourceCleanup(ctx context.Context, store *ResourceStore, logger *slog. } } -func runEndpointCleanup(ctx context.Context, store *EndpointStore, logger *slog.Logger, cfg RetentionConfig, pass *retentionPass) { +func runEndpointCleanup(ctx context.Context, store *EndpointStore, uptime *UptimeDailyStore, logger *slog.Logger, cfg RetentionConfig, pass *retentionPass) { // Clean old check results - cutoff := time.Now().Add(-cfg.CheckResults) - deleted, truncated, err := store.deleteCheckResultsBefore(ctx, cutoff, cfg.batchOpts()) - pass.add(deleted, truncated) - if err != nil { - logger.Error("retention cleanup: check results", "error", err) - } else if deleted > 0 { - logger.Info("retention cleanup: deleted check results", "count", deleted) + if rolledUp(ctx, logger, uptime, "endpoints", cfg.CheckResults, (*UptimeDailyStore).rollupEndpoints) { + cutoff := time.Now().Add(-cfg.CheckResults) + deleted, truncated, err := store.deleteCheckResultsBefore(ctx, cutoff, cfg.batchOpts()) + pass.add(deleted, truncated) + if err != nil { + logger.Error("retention cleanup: check results", "error", err) + } else if deleted > 0 { + logger.Info("retention cleanup: deleted check results", "count", deleted) + } } // Clean inactive endpoints @@ -333,6 +346,34 @@ func runEndpointCleanup(ctx context.Context, store *EndpointStore, logger *slog. } } +// rolledUp writes the daily uptime aggregates of one monitor kind and reports +// whether its raw rows may be purged: a day the rollup could not write stays in +// the raw table until a later pass succeeds. +func rolledUp(ctx context.Context, logger *slog.Logger, uptime *UptimeDailyStore, kind string, rawRetention time.Duration, + rollup func(*UptimeDailyStore, context.Context, time.Time, time.Duration) error) bool { + if uptime == nil { + return true + } + if err := rollup(uptime, ctx, time.Now(), rawRetention); err != nil { + logger.Error("retention cleanup: uptime rollup failed, raw rows kept", "kind", kind, "error", err) + return false + } + return true +} + +func runUptimeDailyCleanup(ctx context.Context, uptime *UptimeDailyStore, logger *slog.Logger, cfg RetentionConfig, pass *retentionPass) { + cutoff := time.Now().Add(-cfg.UptimeDaily) + for _, t := range []uptimeTable{endpointUptimeTable, heartbeatUptimeTable, containerUptimeTable} { + deleted, truncated, err := uptime.deleteBefore(ctx, t, cutoff, cfg.batchOpts()) + pass.add(deleted, truncated) + if err != nil { + logger.Error("retention cleanup: "+t.name, "error", err) + } else if deleted > 0 { + logger.Info("retention cleanup: deleted "+t.name, "count", deleted) + } + } +} + // reclaimSpace returns the pages freed by the pass to the filesystem and keeps // the WAL from carrying the whole reclaim at once. On a database whose header // says auto_vacuum NONE there is nothing to reclaim without a full VACUUM, so diff --git a/internal/store/transform.go b/internal/store/transform.go index 883f8880..cf1edaa9 100644 --- a/internal/store/transform.go +++ b/internal/store/transform.go @@ -140,10 +140,11 @@ func runConversion(ctx context.Context, conn *sql.Conn) error { // Drop any pre-existing agents/enrollment_tokens from the never-deployed // agent migrations (22-23) so the new schema can recreate them cleanly. No-op // on production databases (migrations 1-21 never created these tables). - // instances, cve_evaluations and outbound_heartbeats were created empty by - // later migrations moments ago (none is a legacy table), and uuid_schema - // recreates them below. - for _, t := range []string{"agents", "enrollment_tokens", "instances", "cve_evaluations", "outbound_heartbeats"} { + // instances, cve_evaluations, outbound_heartbeats and the uptime_daily + // tables were created empty by later migrations moments ago (none is a + // legacy table), and uuid_schema recreates them below. + for _, t := range []string{"agents", "enrollment_tokens", "instances", "cve_evaluations", "outbound_heartbeats", + "endpoint_uptime_daily", "heartbeat_uptime_daily", "container_uptime_daily"} { if err := exec("drop unreleased "+t, fmt.Sprintf("DROP TABLE IF EXISTS %q", t)); err != nil { return err } diff --git a/internal/store/uptime_daily.go b/internal/store/uptime_daily.go index 3fa654f3..e5862f4d 100644 --- a/internal/store/uptime_daily.go +++ b/internal/store/uptime_daily.go @@ -8,121 +8,232 @@ import ( "database/sql" "errors" "fmt" + "math" "time" "github.com/kolapsis/maintenant/internal/container" + "github.com/kolapsis/maintenant/internal/uid" ) -// DailyUptime represents a single day's uptime aggregation. +// DailyUptime is the uptime of one monitor over one UTC day. type DailyUptime struct { Date string `json:"date"` UptimePercent *float64 `json:"uptime_percent"` IncidentCount int `json:"incident_count"` } -// UptimeDailyStore provides daily uptime aggregation queries. +// UptimeDailyStore serves per-day uptime: completed days from the daily +// aggregates, the days not aggregated yet computed from the raw rows. type UptimeDailyStore struct { - db *Reader + db *Reader + writer *Writer } // NewUptimeDailyStore creates a new daily uptime store. func NewUptimeDailyStore(d *DB) *UptimeDailyStore { return &UptimeDailyStore{ - db: d.Reader(), + db: d.Reader(), + writer: d.Writer(), } } -// GetEndpointDailyUptime aggregates endpoint check results by UTC day. -// Returns up to `days` days of data, most recent first. -// Days with no checks have UptimePercent = nil. +const ( + maxUptimeDays = 365 + secondsPerDay = 86400 + uptimeDay = 24 * time.Hour + incidentLookback = 24 * time.Hour // how far back the check preceding a day is searched for +) + +type dayUptime struct { + percent float64 + incidents int +} + +// dayUptimes maps the UTC midnight of a day, in epoch seconds, to its uptime. +type dayUptimes map[int64]dayUptime + +// computeDays returns the uptime of monitor id for each day of [from, to) that +// has raw data; from is a UTC midnight. +type computeDays func(ctx context.Context, id string, from, to time.Time) (dayUptimes, error) + +type uptimeTable struct { + name string + column string +} + +var ( + endpointUptimeTable = uptimeTable{name: "endpoint_uptime_daily", column: "endpoint_id"} + heartbeatUptimeTable = uptimeTable{name: "heartbeat_uptime_daily", column: "heartbeat_id"} + containerUptimeTable = uptimeTable{name: "container_uptime_daily", column: "container_id"} +) + +// GetEndpointDailyUptime returns one entry per day, most recent first. func (s *UptimeDailyStore) GetEndpointDailyUptime(ctx context.Context, endpointID string, days int) ([]DailyUptime, error) { + return s.daily(ctx, endpointUptimeTable, endpointID, days, s.endpointDays) +} + +// GetHeartbeatDailyUptime returns one entry per day, most recent first. +func (s *UptimeDailyStore) GetHeartbeatDailyUptime(ctx context.Context, heartbeatID string, days int) ([]DailyUptime, error) { + return s.daily(ctx, heartbeatUptimeTable, heartbeatID, days, s.heartbeatDays) +} + +// GetContainerDailyUptime returns one entry per day, most recent first. +func (s *UptimeDailyStore) GetContainerDailyUptime(ctx context.Context, containerID string, days int) ([]DailyUptime, error) { + until, err := s.containerUntil(ctx, containerID) + if err != nil { + return nil, err + } + return s.daily(ctx, containerUptimeTable, containerID, days, + func(ctx context.Context, id string, from, to time.Time) (dayUptimes, error) { + if until != nil && until.Before(to) { + to = *until + } + return s.containerDays(ctx, id, from, to) + }) +} + +func (s *UptimeDailyStore) daily(ctx context.Context, t uptimeTable, id string, days int, compute computeDays) ([]DailyUptime, error) { days = clampUptimeDays(days) - // Calculate the start of the window (beginning of the day N days ago in UTC). now := time.Now().UTC() - startOfToday := time.Date(now.Year(), now.Month(), now.Day(), 0, 0, 0, 0, time.UTC) - windowStart := startOfToday.AddDate(0, 0, -(days - 1)) - - // Query: aggregate check_results by UTC day. - // the success column is 1 for success, 0 for failure. - // incident_count = number of transitions from success to failure within the day. - rows, err := s.db.QueryContext(ctx, ` - SELECT - date(timestamp, 'unixepoch') AS day, - ROUND(CAST(SUM(success) AS REAL) / COUNT(*) * 100.0, 2) AS uptime_percent, - COUNT(CASE WHEN success = 0 AND prev_success = 1 THEN 1 END) AS incident_count - FROM ( - SELECT - timestamp, - success, - LAG(success) OVER (ORDER BY timestamp) AS prev_success - FROM check_results - WHERE endpoint_id = ? AND timestamp >= ? - ) - GROUP BY day - ORDER BY day DESC - `, endpointID, windowStart.Unix()) + today := startOfUTCDay(now) + windowStart := today.AddDate(0, 0, -(days - 1)) + + stored, err := s.storedDays(ctx, t, id, windowStart, today) + if err != nil { + return nil, err + } + + computeFrom := windowStart + for d := range stored { + if next := time.Unix(d, 0).UTC().Add(uptimeDay); next.After(computeFrom) { + computeFrom = next + } + } + computed, err := compute(ctx, id, computeFrom, now) + if err != nil { + return nil, err + } + + result := make([]DailyUptime, 0, days) + for i := 0; i < days; i++ { + day := today.AddDate(0, 0, -i) + du := DailyUptime{Date: day.Format("2006-01-02")} + v, ok := stored[day.Unix()] + if !ok { + v, ok = computed[day.Unix()] + } + if ok { + pct := v.percent + du.UptimePercent = &pct + du.IncidentCount = v.incidents + } + result = append(result, du) + } + return result, nil +} + +func (s *UptimeDailyStore) storedDays(ctx context.Context, t uptimeTable, id string, from, to time.Time) (dayUptimes, error) { + rows, err := s.db.QueryContext(ctx, + `SELECT day, uptime_percent, incident_count FROM `+t.name+` + WHERE `+t.column+` = ? AND day >= ? AND day < ?`, + id, from.Unix(), to.Unix()) + if err != nil { + return nil, fmt.Errorf("read %s: %w", t.name, err) + } + defer func() { _ = rows.Close() }() + + out := dayUptimes{} + for rows.Next() { + var day int64 + var v dayUptime + if err := rows.Scan(&day, &v.percent, &v.incidents); err != nil { + return nil, fmt.Errorf("scan %s: %w", t.name, err) + } + out[day] = v + } + if err := rows.Err(); err != nil { + return nil, fmt.Errorf("iterate %s: %w", t.name, err) + } + return out, nil +} + +// endpointDays counts successful checks per day; an incident is a failed check +// following a successful one. +func (s *UptimeDailyStore) endpointDays(ctx context.Context, id string, from, to time.Time) (dayUptimes, error) { + rows, err := s.db.QueryContext(ctx, + `SELECT timestamp, success FROM check_results + WHERE endpoint_id = ? AND timestamp >= ? AND timestamp < ? + ORDER BY timestamp`, + id, from.Add(-incidentLookback).Unix(), to.Unix()) if err != nil { return nil, fmt.Errorf("endpoint daily uptime: %w", err) } - defer func(rows *sql.Rows) { - _ = rows.Close() - }(rows) + defer func() { _ = rows.Close() }() - // Build a map of day -> DailyUptime from query results. - dayMap := make(map[string]*DailyUptime) + agg := newCheckDays(from) for rows.Next() { - var du DailyUptime - var uptimePct float64 - if err := rows.Scan(&du.Date, &uptimePct, &du.IncidentCount); err != nil { + var ts int64 + var success int + if err := rows.Scan(&ts, &success); err != nil { return nil, fmt.Errorf("scan endpoint daily uptime: %w", err) } - du.UptimePercent = &uptimePct - dayMap[du.Date] = &du + agg.add(ts, success != 0) } if err := rows.Err(); err != nil { return nil, fmt.Errorf("iterate endpoint daily uptime: %w", err) } + return agg.result(), nil +} - // Generate a full day range, filling gaps with null uptime. - result := make([]DailyUptime, 0, maxUptimeDays) - for i := 0; i < days; i++ { - day := startOfToday.AddDate(0, 0, -i) - dateStr := day.Format("2006-01-02") - if du, ok := dayMap[dateStr]; ok { - result = append(result, *du) - } else { - result = append(result, DailyUptime{Date: dateStr, UptimePercent: nil, IncidentCount: 0}) - } +// heartbeatDays counts successful runs per day. A run is a completion ping: +// a plain success ping, or an exit code ping, up when the code is 0. Start +// pings open a run and say nothing about its outcome. +func (s *UptimeDailyStore) heartbeatDays(ctx context.Context, id string, from, to time.Time) (dayUptimes, error) { + rows, err := s.db.QueryContext(ctx, + `SELECT timestamp, ping_type, exit_code FROM heartbeat_pings + WHERE heartbeat_id = ? AND ping_type IN ('success', 'exit_code') AND timestamp >= ? AND timestamp < ? + ORDER BY timestamp`, + id, from.Add(-incidentLookback).Unix(), to.Unix()) + if err != nil { + return nil, fmt.Errorf("heartbeat daily uptime: %w", err) } + defer func() { _ = rows.Close() }() - return result, nil + agg := newCheckDays(from) + for rows.Next() { + var ts int64 + var pingType string + var exitCode sql.NullInt64 + if err := rows.Scan(&ts, &pingType, &exitCode); err != nil { + return nil, fmt.Errorf("scan heartbeat daily uptime: %w", err) + } + agg.add(ts, pingType == "success" || (exitCode.Valid && exitCode.Int64 == 0)) + } + if err := rows.Err(); err != nil { + return nil, fmt.Errorf("iterate heartbeat daily uptime: %w", err) + } + return agg.result(), nil } -// GetContainerDailyUptime computes a per-day, time-weighted uptime series for a -// container from its state transitions. Unlike endpoints/heartbeats (discrete -// checks), container uptime is the running+healthy fraction of each day. Days -// before the first recorded transition return nil (no data), most recent first. -func (s *UptimeDailyStore) GetContainerDailyUptime(ctx context.Context, containerID string, days int) ([]DailyUptime, error) { - days = clampUptimeDays(days) - - now := time.Now().UTC() - startOfToday := time.Date(now.Year(), now.Month(), now.Day(), 0, 0, 0, 0, time.UTC) - windowStart := startOfToday.AddDate(0, 0, -(days - 1)) +// containerDays weighs each day by the time spent up, starting from the last +// transition before the window. +func (s *UptimeDailyStore) containerDays(ctx context.Context, id string, from, to time.Time) (dayUptimes, error) { + out := dayUptimes{} + if !to.After(from) { + return out, nil + } - // Seed with the last transition before the window so the earliest days know - // the state they began in. transitions := make([]*container.StateTransition, 0) seed, err := scanTransitionRow(s.db.QueryRowContext(ctx, `SELECT `+transitionColumns+` FROM state_transitions WHERE container_id = ? AND timestamp < ? ORDER BY timestamp DESC LIMIT 1`, - containerID, windowStart.Unix(), + id, from.Unix(), )) switch { case err == nil: transitions = append(transitions, seed) case errors.Is(err, sql.ErrNoRows): - // No prior transition; data (if any) starts inside the window. default: return nil, fmt.Errorf("container daily uptime seed: %w", err) } @@ -130,15 +241,13 @@ func (s *UptimeDailyStore) GetContainerDailyUptime(ctx context.Context, containe rows, err := s.db.QueryContext(ctx, `SELECT `+transitionColumns+` FROM state_transitions - WHERE container_id = ? AND timestamp >= ? ORDER BY timestamp ASC`, - containerID, windowStart.Unix(), + WHERE container_id = ? AND timestamp >= ? AND timestamp < ? ORDER BY timestamp ASC`, + id, from.Unix(), to.Unix(), ) if err != nil { return nil, fmt.Errorf("container daily uptime: %w", err) } - defer func(rows *sql.Rows) { - _ = rows.Close() - }(rows) + defer func() { _ = rows.Close() }() for rows.Next() { t, err := scanTransitionRow(rows) if err != nil { @@ -150,49 +259,53 @@ func (s *UptimeDailyStore) GetContainerDailyUptime(ctx context.Context, containe return nil, fmt.Errorf("iterate container daily uptime: %w", err) } - // Determine when data begins: with a seed it predates the window; otherwise - // it starts at the first in-window transition. No transitions => no data. - var dataStart time.Time - hasData := false + if len(transitions) == 0 { + return out, nil + } + dataStart := transitions[0].Timestamp if seeded { - dataStart = windowStart - hasData = true - } else if len(transitions) > 0 { - dataStart = transitions[0].Timestamp - hasData = true + dataStart = from } - result := make([]DailyUptime, 0, maxUptimeDays) - for i := 0; i < days; i++ { - day := startOfToday.AddDate(0, 0, -i) - dateStr := day.Format("2006-01-02") - dayEnd := day.AddDate(0, 0, 1) - if dayEnd.After(now) { - dayEnd = now + for day := from; day.Before(to); day = day.Add(uptimeDay) { + dayEnd := day.Add(uptimeDay) + if dayEnd.After(to) { + dayEnd = to } - - from := day - if from.Before(dataStart) { - from = dataStart + start := day + if start.Before(dataStart) { + start = dataStart } - - if !hasData || !dayEnd.After(from) { - result = append(result, DailyUptime{Date: dateStr, UptimePercent: nil, IncidentCount: 0}) + if !dayEnd.After(start) { continue } - - pct := container.ComputeUptime(transitions, from, dayEnd) - result = append(result, DailyUptime{ - Date: dateStr, - UptimePercent: &pct, - IncidentCount: countContainerIncidents(transitions, from, dayEnd), - }) + out[day.Unix()] = dayUptime{ + percent: container.ComputeUptime(transitions, start, dayEnd), + incidents: countContainerIncidents(transitions, start, dayEnd), + } } + return out, nil +} - return result, nil +// containerUntil returns when an archived container stopped existing, nil +// while it still runs. +func (s *UptimeDailyStore) containerUntil(ctx context.Context, id string) (*time.Time, error) { + var archived int + var archivedAt sql.NullInt64 + err := s.db.QueryRowContext(ctx, `SELECT archived, archived_at FROM containers WHERE id = ?`, id). + Scan(&archived, &archivedAt) + switch { + case errors.Is(err, sql.ErrNoRows): + return nil, nil + case err != nil: + return nil, fmt.Errorf("container daily uptime: %w", err) + case archived == 0 || !archivedAt.Valid: + return nil, nil + } + until := time.Unix(archivedAt.Int64, 0).UTC() + return &until, nil } -// countContainerIncidents counts up->down transitions within [from, to). func countContainerIncidents(transitions []*container.StateTransition, from, to time.Time) int { n := 0 for _, t := range transitions { @@ -210,74 +323,177 @@ func countContainerIncidents(transitions []*container.StateTransition, from, to return n } -// GetHeartbeatDailyUptime aggregates heartbeat pings by UTC day. -// Returns up to `days` days of data, most recent first. -// Days with no pings have UptimePercent = nil. -func (s *UptimeDailyStore) GetHeartbeatDailyUptime(ctx context.Context, heartbeatID string, days int) ([]DailyUptime, error) { - days = clampUptimeDays(days) +// checkDays tallies ordered up/down samples into days. Samples before from only +// tell whether the first sample of the window follows an up one. +type checkDays struct { + from int64 + prevUp bool + hasPrev bool + counts map[int64]*checkCount +} - now := time.Now().UTC() - startOfToday := time.Date(now.Year(), now.Month(), now.Day(), 0, 0, 0, 0, time.UTC) - windowStart := startOfToday.AddDate(0, 0, -(days - 1)) - - // For heartbeat pings, success pings are ping_type='success'. - // We count total pings and success pings per day. - // incident_count = transitions from success to a non-success ping type. - rows, err := s.db.QueryContext(ctx, ` - SELECT - date(timestamp, 'unixepoch') AS day, - ROUND( - CAST(SUM(CASE WHEN ping_type = 'success' THEN 1 ELSE 0 END) AS REAL) - / COUNT(*) * 100.0, 2 - ) AS uptime_percent, - COUNT(CASE WHEN ping_type != 'success' AND prev_type = 'success' THEN 1 END) AS incident_count - FROM ( - SELECT - timestamp, - ping_type, - LAG(ping_type) OVER (ORDER BY timestamp) AS prev_type - FROM heartbeat_pings - WHERE heartbeat_id = ? AND timestamp >= ? - ) - GROUP BY day - ORDER BY day DESC - `, heartbeatID, windowStart.Unix()) +type checkCount struct { + up, total, incidents int +} + +func newCheckDays(from time.Time) *checkDays { + return &checkDays{from: from.Unix(), counts: map[int64]*checkCount{}} +} + +func (c *checkDays) add(ts int64, up bool) { + if ts >= c.from { + day := ts - ts%secondsPerDay + n, ok := c.counts[day] + if !ok { + n = &checkCount{} + c.counts[day] = n + } + n.total++ + switch { + case up: + n.up++ + case c.hasPrev && c.prevUp: + n.incidents++ + } + } + c.prevUp, c.hasPrev = up, true +} + +func (c *checkDays) result() dayUptimes { + out := make(dayUptimes, len(c.counts)) + for day, n := range c.counts { + out[day] = dayUptime{ + percent: math.Round(float64(n.up)/float64(n.total)*10000) / 100, + incidents: n.incidents, + } + } + return out +} + +type uptimeMonitor struct { + id string + until *time.Time +} + +// rollupEndpoints aggregates the completed days of every endpoint. +func (s *UptimeDailyStore) rollupEndpoints(ctx context.Context, now time.Time, rawRetention time.Duration) error { + monitors, err := s.monitors(ctx, `SELECT id, NULL FROM endpoints`) if err != nil { - return nil, fmt.Errorf("heartbeat daily uptime: %w", err) + return err } - defer func(rows *sql.Rows) { - _ = rows.Close() - }(rows) + return s.rollup(ctx, endpointUptimeTable, monitors, now, rawRetention, s.endpointDays) +} - dayMap := make(map[string]*DailyUptime) +// rollupHeartbeats aggregates the completed days of every heartbeat. +func (s *UptimeDailyStore) rollupHeartbeats(ctx context.Context, now time.Time, rawRetention time.Duration) error { + monitors, err := s.monitors(ctx, `SELECT id, NULL FROM heartbeats`) + if err != nil { + return err + } + return s.rollup(ctx, heartbeatUptimeTable, monitors, now, rawRetention, s.heartbeatDays) +} + +// rollupContainers aggregates the completed days of every container, up to its +// archival for an archived one. +func (s *UptimeDailyStore) rollupContainers(ctx context.Context, now time.Time, rawRetention time.Duration) error { + monitors, err := s.monitors(ctx, `SELECT id, CASE WHEN archived = 1 THEN archived_at END FROM containers`) + if err != nil { + return err + } + return s.rollup(ctx, containerUptimeTable, monitors, now, rawRetention, s.containerDays) +} + +func (s *UptimeDailyStore) monitors(ctx context.Context, query string) ([]uptimeMonitor, error) { + rows, err := s.db.QueryContext(ctx, query) + if err != nil { + return nil, fmt.Errorf("list uptime monitors: %w", err) + } + defer func() { _ = rows.Close() }() + + var out []uptimeMonitor for rows.Next() { - var du DailyUptime - var uptimePct float64 - if err := rows.Scan(&du.Date, &uptimePct, &du.IncidentCount); err != nil { - return nil, fmt.Errorf("scan heartbeat daily uptime: %w", err) + var m uptimeMonitor + var until sql.NullInt64 + if err := rows.Scan(&m.id, &until); err != nil { + return nil, fmt.Errorf("scan uptime monitor: %w", err) + } + if until.Valid { + t := time.Unix(until.Int64, 0).UTC() + m.until = &t } - du.UptimePercent = &uptimePct - dayMap[du.Date] = &du + out = append(out, m) } if err := rows.Err(); err != nil { - return nil, fmt.Errorf("iterate heartbeat daily uptime: %w", err) + return nil, fmt.Errorf("iterate uptime monitors: %w", err) } + return out, nil +} - result := make([]DailyUptime, 0, maxUptimeDays) - for i := 0; i < days; i++ { - day := startOfToday.AddDate(0, 0, -i) - dateStr := day.Format("2006-01-02") - if du, ok := dayMap[dateStr]; ok { - result = append(result, *du) - } else { - result = append(result, DailyUptime{Date: dateStr, UptimePercent: nil, IncidentCount: 0}) +// rollup writes the uptime of each completed day from the last aggregated one +// on. It starts no earlier than the first day whose raw rows the purge has not +// touched yet, so a day is never rewritten from partial data. +func (s *UptimeDailyStore) rollup(ctx context.Context, t uptimeTable, monitors []uptimeMonitor, now time.Time, rawRetention time.Duration, compute computeDays) error { + today := startOfUTCDay(now) + from := startOfUTCDay(now.Add(-rawRetention)).Add(uptimeDay) + + var last sql.NullInt64 + if err := s.db.QueryRowContext(ctx, `SELECT MAX(day) FROM `+t.name).Scan(&last); err != nil { + return fmt.Errorf("last day of %s: %w", t.name, err) + } + if last.Valid { + if d := time.Unix(last.Int64, 0).UTC(); d.After(from) { + from = d } } - return result, nil + type aggregate struct { + id string + day int64 + v dayUptime + } + var rows []aggregate + for _, m := range monitors { + to := today + if m.until != nil && m.until.Before(to) { + to = *m.until + } + if !to.After(from) { + continue + } + days, err := compute(ctx, m.id, from, to) + if err != nil { + return err + } + for day, v := range days { + rows = append(rows, aggregate{id: m.id, day: day, v: v}) + } + } + if len(rows) == 0 { + return nil + } + + upsert := `INSERT INTO ` + t.name + ` (id, ` + t.column + `, day, uptime_percent, incident_count) + VALUES (?, ?, ?, ?, ?) + ON CONFLICT(` + t.column + `, day) DO UPDATE SET + uptime_percent = excluded.uptime_percent, incident_count = excluded.incident_count` + return s.writer.Tx(ctx, func(ctx context.Context, tx *Tx) error { + for _, r := range rows { + if _, err := tx.ExecContext(ctx, upsert, uid.New(), r.id, r.day, r.v.percent, r.v.incidents); err != nil { + return fmt.Errorf("write %s: %w", t.name, err) + } + } + return nil + }) +} + +func (s *UptimeDailyStore) deleteBefore(ctx context.Context, t uptimeTable, before time.Time, o batchOpts) (int64, bool, error) { + return deleteRowsBefore(ctx, s.writer, o, t.name, "day", before) } -const maxUptimeDays = 365 +func startOfUTCDay(t time.Time) time.Time { + t = t.UTC() + return time.Date(t.Year(), t.Month(), t.Day(), 0, 0, 0, 0, time.UTC) +} func clampUptimeDays(days int) int { if days <= 0 { diff --git a/internal/store/uptime_daily_test.go b/internal/store/uptime_daily_test.go index 5fb97322..4613761a 100644 --- a/internal/store/uptime_daily_test.go +++ b/internal/store/uptime_daily_test.go @@ -5,419 +5,396 @@ package store import ( "context" - "database/sql" - "log/slog" - "os" "testing" "time" - "github.com/kolapsis/maintenant/internal/uid" - _ "github.com/mattn/go-sqlite3" "github.com/stretchr/testify/assert" "github.com/stretchr/testify/require" + + "github.com/kolapsis/maintenant/internal/endpoint" + "github.com/kolapsis/maintenant/internal/heartbeat" + "github.com/kolapsis/maintenant/internal/uid" ) -// setupTestDB creates an in-memory SQLite database with the required schema for testing. -func setupTestDB(t *testing.T) *DB { +func seedUptimeEndpoint(t *testing.T, db *DB) string { t.Helper() - logger := slog.New(slog.NewTextHandler(os.Stderr, &slog.HandlerOptions{Level: slog.LevelError})) - rawDB, err := sql.Open("sqlite3", ":memory:") - require.NoError(t, err) - t.Cleanup(func() { _ = rawDB.Close() }) - - // Create required tables. - _, err = rawDB.Exec(` - CREATE TABLE IF NOT EXISTS endpoints ( - id TEXT PRIMARY KEY NOT NULL, - container_name TEXT NOT NULL, - label_key TEXT NOT NULL, - external_id TEXT NOT NULL DEFAULT '', - endpoint_type TEXT NOT NULL DEFAULT 'http', - target TEXT NOT NULL DEFAULT '', - status TEXT NOT NULL DEFAULT 'unknown', - alert_state TEXT NOT NULL DEFAULT 'normal', - consecutive_failures INTEGER NOT NULL DEFAULT 0, - consecutive_successes INTEGER NOT NULL DEFAULT 0, - last_check_at INTEGER, - last_response_time_ms INTEGER, - last_http_status INTEGER, - last_error TEXT, - config_json TEXT NOT NULL DEFAULT '{}', - active INTEGER NOT NULL DEFAULT 1, - first_seen_at INTEGER NOT NULL, - last_seen_at INTEGER NOT NULL - ); - CREATE TABLE IF NOT EXISTS check_results ( - id TEXT PRIMARY KEY NOT NULL, - endpoint_id TEXT NOT NULL, - success INTEGER NOT NULL, - response_time_ms INTEGER NOT NULL DEFAULT 0, - http_status INTEGER, - error_message TEXT, - timestamp INTEGER NOT NULL - ); - CREATE TABLE IF NOT EXISTS heartbeats ( - id TEXT PRIMARY KEY NOT NULL, - uuid TEXT NOT NULL, - name TEXT NOT NULL, - status TEXT NOT NULL DEFAULT 'new', - alert_state TEXT NOT NULL DEFAULT 'normal', - interval_seconds INTEGER NOT NULL DEFAULT 300, - grace_seconds INTEGER NOT NULL DEFAULT 60, - last_ping_at INTEGER, - next_deadline_at INTEGER, - current_run_started_at INTEGER, - last_exit_code INTEGER, - last_duration_ms INTEGER, - consecutive_failures INTEGER NOT NULL DEFAULT 0, - consecutive_successes INTEGER NOT NULL DEFAULT 0, - active INTEGER NOT NULL DEFAULT 1, - created_at INTEGER NOT NULL, - updated_at INTEGER NOT NULL - ); - CREATE TABLE IF NOT EXISTS heartbeat_pings ( - id TEXT PRIMARY KEY NOT NULL, - heartbeat_id TEXT NOT NULL, - ping_type TEXT NOT NULL, - exit_code INTEGER, - source_ip TEXT NOT NULL DEFAULT '', - http_method TEXT NOT NULL DEFAULT 'GET', - payload TEXT, - timestamp INTEGER NOT NULL - ); - CREATE TABLE IF NOT EXISTS state_transitions ( - id TEXT PRIMARY KEY NOT NULL, - container_id TEXT NOT NULL, - previous_state TEXT NOT NULL, - new_state TEXT NOT NULL, - previous_health TEXT, - new_health TEXT, - exit_code INTEGER, - log_snippet TEXT, - timestamp INTEGER NOT NULL - ); - `) + id, err := NewEndpointStore(db).UpsertEndpoint(context.Background(), &endpoint.Endpoint{ + ContainerName: "web-" + uid.New(), + LabelKey: "maintenant.endpoint.http", + ExternalID: "ext-web", + EndpointType: endpoint.TypeHTTP, + Target: "http://web:8080", + }) require.NoError(t, err) + return id +} - db := &DB{ - db: rawDB, - logger: logger, - } - return db +func seedUptimeHeartbeat(t *testing.T, db *DB) string { + t.Helper() + id, err := NewHeartbeatStore(db).CreateHeartbeat(context.Background(), &heartbeat.Heartbeat{ + Name: "backup", IntervalSeconds: 300, GraceSeconds: 60, + }) + require.NoError(t, err) + return id } -func insertCheckResult(t *testing.T, db *Reader, endpointID string, success int, ts time.Time) { +func addCheck(t *testing.T, db *DB, endpointID string, success bool, ts time.Time) { t.Helper() - _, err := db.ExecContext(context.Background(), + ok := 0 + if success { + ok = 1 + } + _, err := db.Writer().Exec(context.Background(), `INSERT INTO check_results (id, endpoint_id, success, response_time_ms, timestamp) VALUES (?, ?, ?, 100, ?)`, - uid.New(), endpointID, success, ts.Unix(), - ) + uid.New(), endpointID, ok, ts.Unix()) require.NoError(t, err) } -func insertHeartbeatPing(t *testing.T, db *Reader, heartbeatID string, pingType string, ts time.Time) { +func addPing(t *testing.T, db *DB, heartbeatID, pingType string, exitCode *int, ts time.Time) { t.Helper() - _, err := db.ExecContext(context.Background(), - `INSERT INTO heartbeat_pings (id, heartbeat_id, ping_type, source_ip, http_method, timestamp) VALUES (?, ?, ?, '127.0.0.1', 'GET', ?)`, - uid.New(), heartbeatID, pingType, ts.Unix(), - ) + _, err := db.Writer().Exec(context.Background(), + `INSERT INTO heartbeat_pings (id, heartbeat_id, ping_type, exit_code, source_ip, http_method, timestamp) + VALUES (?, ?, ?, ?, '127.0.0.1', 'GET', ?)`, + uid.New(), heartbeatID, pingType, exitCode, ts.Unix()) require.NoError(t, err) } -func insertTransition(t *testing.T, db *Reader, containerID, prevState, newState string, ts time.Time) { +func addTransition(t *testing.T, db *DB, containerID, prevState, newState string, ts time.Time) { t.Helper() - _, err := db.ExecContext(context.Background(), + _, err := db.Writer().Exec(context.Background(), `INSERT INTO state_transitions (id, container_id, previous_state, new_state, timestamp) VALUES (?, ?, ?, ?, ?)`, - uid.New(), containerID, prevState, newState, ts.Unix(), - ) + uid.New(), containerID, prevState, newState, ts.Unix()) require.NoError(t, err) } -func TestContainerDailyUptime(t *testing.T) { - now := time.Now().UTC() - today := time.Date(now.Year(), now.Month(), now.Day(), 0, 0, 0, 0, time.UTC) +func exitCode(n int) *int { return &n } - t.Run("no transitions returns null days", func(t *testing.T) { - store := NewUptimeDailyStore(setupTestDB(t)) - result, err := store.GetContainerDailyUptime(context.Background(), "c1", 3) +// dayOf returns the entry for day in a most-recent-first series. +func dayOf(t *testing.T, series []DailyUptime, day time.Time) DailyUptime { + t.Helper() + for _, du := range series { + if du.Date == day.Format("2006-01-02") { + return du + } + } + t.Fatalf("day %s missing from the series", day.Format("2006-01-02")) + return DailyUptime{} +} + +func requirePercent(t *testing.T, du DailyUptime, want float64) { + t.Helper() + require.NotNil(t, du.UptimePercent, "day %s has no uptime", du.Date) + assert.InDelta(t, want, *du.UptimePercent, 0.001, "day %s", du.Date) +} + +func TestEndpointDailyUptime(t *testing.T) { + ctx := context.Background() + today := startOfUTCDay(time.Now()) + yesterday := today.AddDate(0, 0, -1) + + t.Run("no checks returns all null days, most recent first", func(t *testing.T) { + db := openTestDB(t) + result, err := NewUptimeDailyStore(db).GetEndpointDailyUptime(ctx, seedUptimeEndpoint(t, db), 3) + require.NoError(t, err) + require.Len(t, result, 3) + assert.Equal(t, today.Format("2006-01-02"), result[0].Date) + assert.Equal(t, yesterday.Format("2006-01-02"), result[1].Date) + for _, du := range result { + assert.Nil(t, du.UptimePercent) + assert.Zero(t, du.IncidentCount) + } + }) + + t.Run("partial uptime with incident", func(t *testing.T) { + db := openTestDB(t) + id := seedUptimeEndpoint(t, db) + for i := 0; i < 4; i++ { + addCheck(t, db, id, true, yesterday.Add(time.Duration(i)*time.Hour)) + } + addCheck(t, db, id, false, yesterday.Add(4*time.Hour)) + + result, err := NewUptimeDailyStore(db).GetEndpointDailyUptime(ctx, id, 2) + require.NoError(t, err) + requirePercent(t, result[1], 80) + assert.Equal(t, 1, result[1].IncidentCount) + }) + + t.Run("a day opening on a failure after a success counts an incident", func(t *testing.T) { + db := openTestDB(t) + id := seedUptimeEndpoint(t, db) + addCheck(t, db, id, true, yesterday.Add(-time.Hour)) + addCheck(t, db, id, false, yesterday.Add(time.Minute)) + + result, err := NewUptimeDailyStore(db).GetEndpointDailyUptime(ctx, id, 2) + require.NoError(t, err) + requirePercent(t, result[1], 0) + assert.Equal(t, 1, result[1].IncidentCount) + }) + + t.Run("gap days stay null", func(t *testing.T) { + db := openTestDB(t) + id := seedUptimeEndpoint(t, db) + addCheck(t, db, id, false, today.AddDate(0, 0, -3).Add(5*time.Hour)) + addCheck(t, db, id, true, yesterday.Add(time.Hour)) + + result, err := NewUptimeDailyStore(db).GetEndpointDailyUptime(ctx, id, 4) + require.NoError(t, err) + requirePercent(t, result[1], 100) + assert.Nil(t, result[2].UptimePercent) + requirePercent(t, result[3], 0) + }) + + t.Run("days default to 90 and cap at 365", func(t *testing.T) { + db := openTestDB(t) + id := seedUptimeEndpoint(t, db) + store := NewUptimeDailyStore(db) + + result, err := store.GetEndpointDailyUptime(ctx, id, 0) + require.NoError(t, err) + assert.Len(t, result, 90) + + result, err = store.GetEndpointDailyUptime(ctx, id, 500) + require.NoError(t, err) + assert.Len(t, result, 365) + }) +} + +func TestHeartbeatDailyUptime(t *testing.T) { + ctx := context.Background() + yesterday := startOfUTCDay(time.Now()).AddDate(0, 0, -1) + + t.Run("no pings returns null days", func(t *testing.T) { + db := openTestDB(t) + result, err := NewUptimeDailyStore(db).GetHeartbeatDailyUptime(ctx, seedUptimeHeartbeat(t, db), 3) require.NoError(t, err) - assert.Len(t, result, 3) + require.Len(t, result, 3) for _, du := range result { - assert.Nil(t, du.UptimePercent, "day %s should be null with no data", du.Date) - assert.Equal(t, 0, du.IncidentCount) + assert.Nil(t, du.UptimePercent) } }) - t.Run("running since before the window is 100% today", func(t *testing.T) { - d := setupTestDB(t) - insertTransition(t, d.Reader(), "c1", "created", "running", today.AddDate(0, 0, -5)) - result, err := NewUptimeDailyStore(d).GetContainerDailyUptime(context.Background(), "c1", 1) + t.Run("a run reported with start then exit code 0 is up", func(t *testing.T) { + db := openTestDB(t) + id := seedUptimeHeartbeat(t, db) + for i := 0; i < 3; i++ { + run := yesterday.Add(time.Duration(i) * time.Hour) + addPing(t, db, id, "start", nil, run) + addPing(t, db, id, "exit_code", exitCode(0), run.Add(time.Minute)) + } + + result, err := NewUptimeDailyStore(db).GetHeartbeatDailyUptime(ctx, id, 2) require.NoError(t, err) - require.Len(t, result, 1) - require.NotNil(t, result[0].UptimePercent) - assert.Equal(t, 100.0, *result[0].UptimePercent) - assert.Equal(t, 0, result[0].IncidentCount) + requirePercent(t, result[1], 100) + assert.Zero(t, result[1].IncidentCount) + }) + + t.Run("a failing exit code after a success is an incident", func(t *testing.T) { + db := openTestDB(t) + id := seedUptimeHeartbeat(t, db) + addPing(t, db, id, "success", nil, yesterday.Add(1*time.Hour)) + addPing(t, db, id, "success", nil, yesterday.Add(2*time.Hour)) + addPing(t, db, id, "exit_code", exitCode(0), yesterday.Add(3*time.Hour)) + addPing(t, db, id, "exit_code", exitCode(2), yesterday.Add(4*time.Hour)) + + result, err := NewUptimeDailyStore(db).GetHeartbeatDailyUptime(ctx, id, 2) + require.NoError(t, err) + requirePercent(t, result[1], 75) + assert.Equal(t, 1, result[1].IncidentCount) + }) +} + +func TestContainerDailyUptime(t *testing.T) { + ctx := context.Background() + today := startOfUTCDay(time.Now()) + + t.Run("no transitions returns null days", func(t *testing.T) { + db := openTestDB(t) + cid := seedHostContainer(t, NewContainerStore(db), "ext-none", "") + result, err := NewUptimeDailyStore(db).GetContainerDailyUptime(ctx, cid, 3) + require.NoError(t, err) + require.Len(t, result, 3) + for _, du := range result { + assert.Nil(t, du.UptimePercent) + } }) t.Run("full past day with a down period is time-weighted", func(t *testing.T) { - d := setupTestDB(t) - // Running well before the window so every day is seeded as up. - insertTransition(t, d.Reader(), "c1", "created", "running", today.AddDate(0, 0, -10)) - // Two days ago: down from +8h to +16h (8h of a full 24h day -> 66.66% up). + db := openTestDB(t) + cid := seedHostContainer(t, NewContainerStore(db), "ext-weighted", "") + addTransition(t, db, cid, "created", "running", today.AddDate(0, 0, -10)) twoDaysAgo := today.AddDate(0, 0, -2) - insertTransition(t, d.Reader(), "c1", "running", "exited", twoDaysAgo.Add(8*time.Hour)) - insertTransition(t, d.Reader(), "c1", "exited", "running", twoDaysAgo.Add(16*time.Hour)) + addTransition(t, db, cid, "running", "exited", twoDaysAgo.Add(8*time.Hour)) + addTransition(t, db, cid, "exited", "running", twoDaysAgo.Add(16*time.Hour)) - result, err := NewUptimeDailyStore(d).GetContainerDailyUptime(context.Background(), "c1", 3) + result, err := NewUptimeDailyStore(db).GetContainerDailyUptime(ctx, cid, 3) require.NoError(t, err) - require.Len(t, result, 3) - // Most recent first: [0]=today, [1]=yesterday, [2]=two days ago. - require.NotNil(t, result[2].UptimePercent) - assert.Equal(t, 66.66, *result[2].UptimePercent) + requirePercent(t, result[2], 66.66) assert.Equal(t, 1, result[2].IncidentCount) - require.NotNil(t, result[1].UptimePercent) - assert.Equal(t, 100.0, *result[1].UptimePercent) + requirePercent(t, result[1], 100) }) - t.Run("days clamped from 0 to 90", func(t *testing.T) { - store := NewUptimeDailyStore(setupTestDB(t)) - result, err := store.GetContainerDailyUptime(context.Background(), "c1", 0) + t.Run("days after its archival have no uptime", func(t *testing.T) { + db := openTestDB(t) + cs := NewContainerStore(db) + cid := seedHostContainer(t, cs, "ext-archived", "") + addTransition(t, db, cid, "created", "running", today.AddDate(0, 0, -5)) + addTransition(t, db, cid, "running", "exited", today.AddDate(0, 0, -3).Add(12*time.Hour)) + require.NoError(t, cs.ArchiveContainer(ctx, cid, today.AddDate(0, 0, -3).Add(13*time.Hour))) + + result, err := NewUptimeDailyStore(db).GetContainerDailyUptime(ctx, cid, 5) require.NoError(t, err) - assert.Len(t, result, 90) + requirePercent(t, result[4], 100) + require.NotNil(t, result[3].UptimePercent) + assert.Nil(t, result[2].UptimePercent, "the container no longer existed") + assert.Nil(t, result[0].UptimePercent) }) } -func TestEndpointDailyUptime(t *testing.T) { - now := time.Now().UTC() - today := time.Date(now.Year(), now.Month(), now.Day(), 0, 0, 0, 0, time.UTC) - - tests := []struct { - name string - endpointID string - days int - setup func(t *testing.T, db *Reader) - wantLen int - checkFirstDay func(t *testing.T, du DailyUptime) - checkNullDays bool // expect null uptime for days with no data - }{ - { - name: "no checks returns all null days", - endpointID: "1", - days: 3, - setup: func(t *testing.T, db *Reader) {}, - wantLen: 3, - checkFirstDay: func(t *testing.T, du DailyUptime) { - assert.Equal(t, today.Format("2006-01-02"), du.Date) - assert.Nil(t, du.UptimePercent, "no checks should yield null uptime") - assert.Equal(t, 0, du.IncidentCount) - }, - checkNullDays: true, - }, - { - name: "100% uptime day", - endpointID: "1", - days: 1, - setup: func(t *testing.T, db *Reader) { - for i := 0; i < 10; i++ { - insertCheckResult(t, db, "1", 1, today.Add(time.Duration(i)*time.Hour)) - } - }, - wantLen: 1, - checkFirstDay: func(t *testing.T, du DailyUptime) { - require.NotNil(t, du.UptimePercent) - assert.Equal(t, 100.0, *du.UptimePercent) - assert.Equal(t, 0, du.IncidentCount) - }, - }, - { - name: "0% uptime day", - endpointID: "2", - days: 1, - setup: func(t *testing.T, db *Reader) { - for i := 0; i < 5; i++ { - insertCheckResult(t, db, "2", 0, today.Add(time.Duration(i)*time.Hour)) - } - }, - wantLen: 1, - checkFirstDay: func(t *testing.T, du DailyUptime) { - require.NotNil(t, du.UptimePercent) - assert.Equal(t, 0.0, *du.UptimePercent) - }, - }, - { - name: "partial uptime with incident", - endpointID: "3", - days: 1, - setup: func(t *testing.T, db *Reader) { - // 4 success, then 1 failure = 80% uptime, 1 incident - for i := 0; i < 4; i++ { - insertCheckResult(t, db, "3", 1, today.Add(time.Duration(i)*time.Hour)) - } - insertCheckResult(t, db, "3", 0, today.Add(4*time.Hour)) - }, - wantLen: 1, - checkFirstDay: func(t *testing.T, du DailyUptime) { - require.NotNil(t, du.UptimePercent) - assert.Equal(t, 80.0, *du.UptimePercent) - assert.Equal(t, 1, du.IncidentCount) - }, - }, - { - name: "multi-day with gap", - endpointID: "4", - days: 3, - setup: func(t *testing.T, db *Reader) { - // Today: 2 checks both success - insertCheckResult(t, db, "4", 1, today.Add(1*time.Hour)) - insertCheckResult(t, db, "4", 1, today.Add(2*time.Hour)) - // Yesterday: no checks (should be null) - // Day before: 1 check, failure - twoDaysAgo := today.AddDate(0, 0, -2) - insertCheckResult(t, db, "4", 0, twoDaysAgo.Add(5*time.Hour)) - }, - wantLen: 3, - checkFirstDay: func(t *testing.T, du DailyUptime) { - // Most recent first = today - assert.Equal(t, today.Format("2006-01-02"), du.Date) - require.NotNil(t, du.UptimePercent) - assert.Equal(t, 100.0, *du.UptimePercent) - }, - }, - { - name: "default days clamped from 0 to 90", - endpointID: "1", - days: 0, - setup: func(t *testing.T, db *Reader) {}, - wantLen: 90, - }, - { - name: "max days clamped to 365", - endpointID: "1", - days: 500, - setup: func(t *testing.T, db *Reader) {}, - wantLen: 365, - }, +// A completed day is written to the daily aggregate before its raw rows are +// purged, so the 90-day bars outlive the 30-day raw retention. +func TestUptimeDaily_SurvivesRawPurge(t *testing.T) { + db := openTestDB(t) + ctx := context.Background() + logger := testLogger() + today := startOfUTCDay(time.Now()) + day := today.AddDate(0, 0, -3) + + epID := seedUptimeEndpoint(t, db) + for i := 0; i < 3; i++ { + addCheck(t, db, epID, true, day.Add(time.Duration(i)*time.Hour)) } + addCheck(t, db, epID, false, day.Add(3*time.Hour)) - for _, tt := range tests { - t.Run(tt.name, func(t *testing.T) { - d := setupTestDB(t) - store := NewUptimeDailyStore(d) - tt.setup(t, d.Reader()) - - result, err := store.GetEndpointDailyUptime(context.Background(), tt.endpointID, tt.days) - require.NoError(t, err) - assert.Len(t, result, tt.wantLen) - - if tt.checkFirstDay != nil && len(result) > 0 { - tt.checkFirstDay(t, result[0]) - } - - if tt.checkNullDays { - for _, du := range result { - assert.Nil(t, du.UptimePercent, "day %s should have null uptime", du.Date) - } - } - - // Verify ordering: the most recent first. - if len(result) > 1 { - assert.GreaterOrEqual(t, result[0].Date, result[1].Date, "days should be ordered most recent first") - } - }) + hbID := seedUptimeHeartbeat(t, db) + addPing(t, db, hbID, "success", nil, day.Add(time.Hour)) + addPing(t, db, hbID, "exit_code", exitCode(1), day.Add(2*time.Hour)) + + cs := NewContainerStore(db) + cid := seedHostContainer(t, cs, "ext-purge", "") + addTransition(t, db, cid, "created", "running", day) + addTransition(t, db, cid, "running", "exited", day.Add(18*time.Hour)) + addTransition(t, db, cid, "exited", "running", today.AddDate(0, 0, -2)) + + uptime := NewUptimeDailyStore(db) + eps, hbs := NewEndpointStore(db), NewHeartbeatStore(db) + pass := func(raw time.Duration) { + cfg := RetentionConfig{CheckResults: raw, HeartbeatPings: raw, Transitions: raw}.withDefaults(logger) + var p retentionPass + runCleanup(ctx, cs, uptime, logger, cfg, &p) + runEndpointCleanup(ctx, eps, uptime, logger, cfg, &p) + runHeartbeatCleanup(ctx, hbs, uptime, logger, cfg, &p) } + + pass(30 * 24 * time.Hour) + pass(time.Hour) + + assert.Zero(t, countTableRows(t, db, "check_results"), "the raw checks are gone") + assert.Zero(t, countTableRows(t, db, "heartbeat_pings"), "the raw pings are gone") + + ep, err := uptime.GetEndpointDailyUptime(ctx, epID, 7) + require.NoError(t, err) + requirePercent(t, dayOf(t, ep, day), 75) + assert.Equal(t, 1, dayOf(t, ep, day).IncidentCount) + + hb, err := uptime.GetHeartbeatDailyUptime(ctx, hbID, 7) + require.NoError(t, err) + requirePercent(t, dayOf(t, hb, day), 50) + + ct, err := uptime.GetContainerDailyUptime(ctx, cid, 7) + require.NoError(t, err) + requirePercent(t, dayOf(t, ct, day), 75) + assert.Equal(t, 1, dayOf(t, ct, day).IncidentCount) + requirePercent(t, dayOf(t, ct, today), 100) } -func TestHeartbeatDailyUptime(t *testing.T) { - now := time.Now().UTC() - today := time.Date(now.Year(), now.Month(), now.Day(), 0, 0, 0, 0, time.UTC) - - tests := []struct { - name string - heartbeatID string - days int - setup func(t *testing.T, db *Reader) - wantLen int - checkFirstDay func(t *testing.T, du DailyUptime) - checkNullDays bool - }{ - { - name: "no pings returns null days", - heartbeatID: "1", - days: 3, - setup: func(t *testing.T, db *Reader) {}, - wantLen: 3, - checkFirstDay: func(t *testing.T, du DailyUptime) { - assert.Nil(t, du.UptimePercent) - assert.Equal(t, 0, du.IncidentCount) - }, - checkNullDays: true, - }, - { - name: "all success pings = 100%", - heartbeatID: "1", - days: 1, - setup: func(t *testing.T, db *Reader) { - for i := 0; i < 6; i++ { - insertHeartbeatPing(t, db, "1", "success", today.Add(time.Duration(i)*time.Hour)) - } - }, - wantLen: 1, - checkFirstDay: func(t *testing.T, du DailyUptime) { - require.NotNil(t, du.UptimePercent) - assert.Equal(t, 100.0, *du.UptimePercent) - assert.Equal(t, 0, du.IncidentCount) - }, - }, - { - name: "mixed pings with exit_code type", - heartbeatID: "2", - days: 1, - setup: func(t *testing.T, db *Reader) { - // 3 successes + 1 exit_code (not success) = 75% - insertHeartbeatPing(t, db, "2", "success", today.Add(1*time.Hour)) - insertHeartbeatPing(t, db, "2", "success", today.Add(2*time.Hour)) - insertHeartbeatPing(t, db, "2", "success", today.Add(3*time.Hour)) - insertHeartbeatPing(t, db, "2", "exit_code", today.Add(4*time.Hour)) - }, - wantLen: 1, - checkFirstDay: func(t *testing.T, du DailyUptime) { - require.NotNil(t, du.UptimePercent) - assert.Equal(t, 75.0, *du.UptimePercent) - assert.Equal(t, 1, du.IncidentCount) // success->exit_code transition - }, - }, - { - name: "90 day window default", - heartbeatID: "1", - days: 0, - setup: func(t *testing.T, db *Reader) {}, - wantLen: 90, - }, - } +// A check replayed by an agent after the day was aggregated still lands in it. +func TestUptimeRollup_RewritesTheLastAggregatedDay(t *testing.T) { + db := openTestDB(t) + ctx := context.Background() + now := time.Now() + yesterday := startOfUTCDay(now).AddDate(0, 0, -1) + uptime := NewUptimeDailyStore(db) + + id := seedUptimeEndpoint(t, db) + addCheck(t, db, id, true, yesterday.Add(time.Hour)) + require.NoError(t, uptime.rollupEndpoints(ctx, now, 30*24*time.Hour)) + + addCheck(t, db, id, false, yesterday.Add(2*time.Hour)) + require.NoError(t, uptime.rollupEndpoints(ctx, now, 30*24*time.Hour)) + + stored, err := uptime.storedDays(ctx, endpointUptimeTable, id, yesterday, yesterday.Add(uptimeDay)) + require.NoError(t, err) + require.Contains(t, stored, yesterday.Unix()) + assert.InDelta(t, 50, stored[yesterday.Unix()].percent, 0.001) +} + +// Once the purge has started eating a day, the rollup must not rewrite it from +// what is left. +func TestUptimeRollup_LeavesAPartlyPurgedDayAlone(t *testing.T) { + db := openTestDB(t) + ctx := context.Background() + now := time.Now() + day := startOfUTCDay(now).AddDate(0, 0, -5) + uptime := NewUptimeDailyStore(db) + + id := seedUptimeEndpoint(t, db) + addCheck(t, db, id, true, day.Add(time.Hour)) + addCheck(t, db, id, false, day.Add(2*time.Hour)) + require.NoError(t, uptime.rollupEndpoints(ctx, now, 30*24*time.Hour)) + + _, err := db.Writer().Exec(ctx, `DELETE FROM check_results WHERE endpoint_id = ? AND success = 0`, id) + require.NoError(t, err) + require.NoError(t, uptime.rollupEndpoints(ctx, now, 5*24*time.Hour)) + + stored, err := uptime.storedDays(ctx, endpointUptimeTable, id, day, day.Add(uptimeDay)) + require.NoError(t, err) + assert.InDelta(t, 50, stored[day.Unix()].percent, 0.001) +} - for _, tt := range tests { - t.Run(tt.name, func(t *testing.T) { - d := setupTestDB(t) - store := NewUptimeDailyStore(d) - tt.setup(t, d.Reader()) - - result, err := store.GetHeartbeatDailyUptime(context.Background(), tt.heartbeatID, tt.days) - require.NoError(t, err) - assert.Len(t, result, tt.wantLen) - - if tt.checkFirstDay != nil && len(result) > 0 { - tt.checkFirstDay(t, result[0]) - } - - if tt.checkNullDays { - for _, du := range result { - assert.Nil(t, du.UptimePercent, "day %s should have null uptime", du.Date) - } - } - - if len(result) > 1 { - assert.GreaterOrEqual(t, result[0].Date, result[1].Date, "days should be ordered most recent first") - } - }) +func TestUptimeDailyCleanup_KeepsAYear(t *testing.T) { + db := openTestDB(t) + ctx := context.Background() + today := startOfUTCDay(time.Now()) + id := seedUptimeEndpoint(t, db) + + for _, day := range []time.Time{today.AddDate(0, 0, -364), today.AddDate(0, 0, -400)} { + _, err := db.Writer().Exec(ctx, + `INSERT INTO endpoint_uptime_daily (id, endpoint_id, day, uptime_percent, incident_count) VALUES (?, ?, ?, 99.5, 0)`, + uid.New(), id, day.Unix()) + require.NoError(t, err) } + + var p retentionPass + runUptimeDailyCleanup(ctx, NewUptimeDailyStore(db), testLogger(), RetentionConfig{}.withDefaults(testLogger()), &p) + + assert.Equal(t, int64(1), p.deleted) + result, err := NewUptimeDailyStore(db).GetEndpointDailyUptime(ctx, id, 365) + require.NoError(t, err) + requirePercent(t, dayOf(t, result, today.AddDate(0, 0, -364)), 99.5) +} + +// A container that has not changed state for longer than the transition +// retention still knows the state it is in. +func TestDeleteTransitionsBefore_KeepsTheLatestOfEachContainer(t *testing.T) { + db := openTestDB(t) + ctx := context.Background() + cs := NewContainerStore(db) + old := time.Now().AddDate(0, 0, -200) + + stable := seedHostContainer(t, cs, "ext-stable", "") + addTransition(t, db, stable, "created", "running", old) + + flappy := seedHostContainer(t, cs, "ext-flappy", "") + addTransition(t, db, flappy, "created", "running", old) + addTransition(t, db, flappy, "running", "exited", old.Add(time.Hour)) + addTransition(t, db, flappy, "exited", "running", time.Now().Add(-time.Hour)) + + deleted, err := cs.DeleteTransitionsBefore(ctx, time.Now().AddDate(0, 0, -90), 1000) + require.NoError(t, err) + assert.Equal(t, int64(2), deleted) + assert.Equal(t, 2, countTableRows(t, db, "state_transitions")) + + result, err := NewUptimeDailyStore(db).GetContainerDailyUptime(ctx, stable, 1) + require.NoError(t, err) + requirePercent(t, result[0], 100) } diff --git a/internal/store/uuid_schema.sql b/internal/store/uuid_schema.sql index 61521aea..ade96d7c 100644 --- a/internal/store/uuid_schema.sql +++ b/internal/store/uuid_schema.sql @@ -116,6 +116,15 @@ CREATE TABLE state_transitions ( CREATE INDEX idx_transition_container_time ON state_transitions(container_id, timestamp DESC); CREATE INDEX idx_transition_timestamp ON state_transitions(timestamp); +CREATE TABLE container_uptime_daily ( + id TEXT PRIMARY KEY NOT NULL, + container_id TEXT NOT NULL REFERENCES containers(id) ON DELETE CASCADE, + day BIGINT NOT NULL, -- UTC midnight, epoch seconds + uptime_percent REAL NOT NULL, + incident_count INTEGER NOT NULL, + UNIQUE(container_id, day) +); + CREATE TABLE resource_snapshots ( id TEXT PRIMARY KEY NOT NULL, container_id TEXT NOT NULL REFERENCES containers(id) ON DELETE CASCADE, @@ -221,6 +230,15 @@ CREATE TABLE check_results ( CREATE INDEX idx_check_endpoint_time ON check_results(endpoint_id, timestamp DESC); CREATE INDEX idx_check_timestamp ON check_results(timestamp); +CREATE TABLE endpoint_uptime_daily ( + id TEXT PRIMARY KEY NOT NULL, + endpoint_id TEXT NOT NULL REFERENCES endpoints(id) ON DELETE CASCADE, + day BIGINT NOT NULL, -- UTC midnight, epoch seconds + uptime_percent REAL NOT NULL, + incident_count INTEGER NOT NULL, + UNIQUE(endpoint_id, day) +); + -- ========================================================= cert monitors ===== CREATE TABLE cert_monitors ( id TEXT PRIMARY KEY NOT NULL, -- uid.CertMonitor(agent,host,port[,server_name]) or minted (standalone) @@ -334,6 +352,15 @@ CREATE TABLE heartbeat_pings ( CREATE INDEX idx_hb_ping_heartbeat_time ON heartbeat_pings(heartbeat_id, timestamp DESC); CREATE INDEX idx_hb_ping_timestamp ON heartbeat_pings(timestamp); +CREATE TABLE heartbeat_uptime_daily ( + id TEXT PRIMARY KEY NOT NULL, + heartbeat_id TEXT NOT NULL REFERENCES heartbeats(id) ON DELETE CASCADE, + day BIGINT NOT NULL, -- UTC midnight, epoch seconds + uptime_percent REAL NOT NULL, + incident_count INTEGER NOT NULL, + UNIQUE(heartbeat_id, day) +); + CREATE TABLE heartbeat_executions ( id TEXT PRIMARY KEY NOT NULL, heartbeat_id TEXT NOT NULL REFERENCES heartbeats(id) ON DELETE CASCADE, diff --git a/internal/swarm/labels.go b/internal/swarm/labels.go index 4cbf373c..32821a7b 100644 --- a/internal/swarm/labels.go +++ b/internal/swarm/labels.go @@ -3,58 +3,4 @@ package swarm -import ( - "strconv" - - cmodel "github.com/kolapsis/maintenant/internal/container" -) - -const ( - // Docker Swarm built-in labels. - labelStackNamespace = "com.docker.stack.namespace" - labelSwarmServiceID = "com.docker.swarm.service.id" - - // Maintenant labels (applied at service level in Swarm). - labelMaintGroup = "maintenant.group" - labelMaintIgnore = "maintenant.ignore" - labelMaintSeverity = "maintenant.alert.severity" - labelMaintThreshold = "maintenant.alert.restart_threshold" -) - -// IsSwarmManaged returns true if the container has a Swarm service ID label. -func IsSwarmManaged(labels map[string]string) bool { - _, ok := labels[labelSwarmServiceID] - return ok -} - -// StackName extracts the stack namespace from labels. -func StackName(labels map[string]string) string { - return labels[labelStackNamespace] -} - -// ApplyServiceLabels maps Swarm service-level labels to Container model fields. -// This applies maintenant.* labels from the service definition to the container. -func ApplyServiceLabels(c *cmodel.Container, serviceLabels map[string]string) { - if v, ok := serviceLabels[labelMaintGroup]; ok && v != "" { - c.CustomGroup = v - } - if v, ok := serviceLabels[labelMaintIgnore]; ok && (v == "true" || v == "1") { - c.IsIgnored = true - } - if v, ok := serviceLabels[labelMaintSeverity]; ok { - switch cmodel.AlertSeverity(v) { - case cmodel.SeverityCritical, cmodel.SeverityWarning, cmodel.SeverityInfo: - c.AlertSeverity = cmodel.AlertSeverity(v) - } - } - if v, ok := serviceLabels[labelMaintThreshold]; ok { - if n, err := strconv.Atoi(v); err == nil && n > 0 { - c.RestartThreshold = n - } - } - - // Stack grouping via com.docker.stack.namespace - if stack := serviceLabels[labelStackNamespace]; stack != "" { - c.OrchestrationGroup = stack - } -} +const labelStackNamespace = "com.docker.stack.namespace" diff --git a/internal/swarm/service.go b/internal/swarm/service.go index 5fac39f5..1775e08a 100644 --- a/internal/swarm/service.go +++ b/internal/swarm/service.go @@ -8,11 +8,8 @@ import ( "fmt" "log/slog" "sync" - "time" "github.com/moby/moby/api/types/swarm" - - cmodel "github.com/kolapsis/maintenant/internal/container" ) // ServiceClient abstracts Docker SDK calls needed for Swarm service discovery. @@ -52,38 +49,25 @@ func (sd *ServiceDiscovery) SetNetworkResolver(resolver NetworkResolver) { sd.networkResolver = resolver } -// DiscoverAll discovers all Swarm services and maps them to Container models. -// Returns containers representing Swarm tasks with Swarm fields populated. -func (sd *ServiceDiscovery) DiscoverAll(ctx context.Context) ([]*cmodel.Container, []*SwarmService, error) { +// DiscoverAll discovers all Swarm services and refreshes the service cache. +func (sd *ServiceDiscovery) DiscoverAll(ctx context.Context) ([]*SwarmService, error) { services, err := sd.client.ServiceList(ctx) if err != nil { - return nil, nil, fmt.Errorf("discover services: %w", err) + return nil, fmt.Errorf("discover services: %w", err) } tasks, err := sd.client.TaskList(ctx) if err != nil { - return nil, nil, fmt.Errorf("discover tasks: %w", err) - } - - // Build node hostname map for task placement. - nodeHostnames := make(map[string]string) - nodes, err := sd.client.NodeList(ctx) - if err != nil { - sd.logger.Warn("failed to list nodes for hostname resolution", "error", err) - } else { - for _, n := range nodes { - nodeHostnames[n.ID] = n.Description.Hostname - } + return nil, fmt.Errorf("discover tasks: %w", err) } - // Group tasks by service ID. - tasksByService := make(map[string][]swarm.Task) + runningByService := make(map[string]int) for _, t := range tasks { - tasksByService[t.ServiceID] = append(tasksByService[t.ServiceID], t) + if t.Status.State == swarm.TaskStateRunning { + runningByService[t.ServiceID]++ + } } - now := time.Now() - var containers []*cmodel.Container swarmServices := make([]*SwarmService, 0, len(services)) sd.mu.Lock() @@ -94,37 +78,16 @@ func (sd *ServiceDiscovery) DiscoverAll(ctx context.Context) ([]*cmodel.Containe for _, svc := range services { ss := mapService(svc) - - serviceTasks := tasksByService[svc.ID] - runningCount := 0 - for _, t := range serviceTasks { - if t.Status.State == swarm.TaskStateRunning { - runningCount++ - } - } - ss.RunningReplicas = runningCount + ss.RunningReplicas = runningByService[svc.ID] sd.resolveNetworks(ctx, ss) sd.services[svc.ID] = ss swarmServices = append(swarmServices, ss) - - // Map tasks to containers. - for _, t := range serviceTasks { - // Only map tasks with an active desired state. - if t.DesiredState != swarm.TaskStateRunning && t.DesiredState != swarm.TaskStateShutdown { - continue - } - - c := mapTaskToContainer(svc, t, ss, nodeHostnames, now) - containers = append(containers, c) - } } - sd.logger.Info("discovered Swarm services", - "services", len(services), - "tasks", len(containers)) + sd.logger.Info("discovered Swarm services", "services", len(services)) - return containers, swarmServices, nil + return swarmServices, nil } // GetService returns a cached Swarm service by ID. @@ -312,72 +275,3 @@ func mapService(svc swarm.Service) *SwarmService { return ss } - -func mapTaskToContainer(svc swarm.Service, task swarm.Task, ss *SwarmService, nodeHostnames map[string]string, now time.Time) *cmodel.Container { - containerID := "" - if task.Status.ContainerStatus != nil { - containerID = task.Status.ContainerStatus.ContainerID - } - - name := svc.Spec.Name - if task.Slot > 0 { - name = fmt.Sprintf("%s.%d", svc.Spec.Name, task.Slot) - } - - state := mapTaskState(task.Status.State) - readyCount := 0 - if state == cmodel.StateRunning { - readyCount = 1 - } - - c := &cmodel.Container{ - ExternalID: containerID, - Name: name, - Image: ss.Image, - State: state, - RuntimeType: "docker", - ControllerKind: "swarm-service", - OrchestrationUnit: svc.Spec.Name, - PodCount: 1, - ReadyCount: readyCount, - AlertSeverity: cmodel.SeverityWarning, - RestartThreshold: 3, - FirstSeenAt: now, - LastStateChangeAt: task.Status.Timestamp, - SwarmServiceID: svc.ID, - SwarmServiceName: svc.Spec.Name, - SwarmServiceMode: ss.Mode, - SwarmNodeID: task.NodeID, - SwarmTaskSlot: task.Slot, - SwarmDesiredReplicas: ss.DesiredReplicas, - } - - // Set error detail from task errors. - if task.Status.Err != "" { - c.ErrorDetail = task.Status.Err - } - - // Apply service-level labels. - ApplyServiceLabels(c, svc.Spec.Labels) - - return c -} - -func mapTaskState(state swarm.TaskState) cmodel.ContainerState { - switch state { - case swarm.TaskStateRunning: - return cmodel.StateRunning - case swarm.TaskStateComplete: - return cmodel.StateCompleted - case swarm.TaskStateFailed, swarm.TaskStateRejected: - return cmodel.StateExited - case swarm.TaskStateShutdown: - return cmodel.StateExited - case swarm.TaskStateNew, swarm.TaskStatePending, swarm.TaskStateAssigned, - swarm.TaskStateAccepted, swarm.TaskStatePreparing, swarm.TaskStateStarting, - swarm.TaskStateReady: - return cmodel.StateCreated - default: - return cmodel.StateCreated - } -} diff --git a/internal/swarm/snapshot.go b/internal/swarm/snapshot.go index 816c26d3..fb928280 100644 --- a/internal/swarm/snapshot.go +++ b/internal/swarm/snapshot.go @@ -16,7 +16,7 @@ import ( // ingest service under the LocalAgent id). disc supplies services and tasks; // client supplies nodes. func SnapshotFromClient(ctx context.Context, disc *ServiceDiscovery, client ServiceClient) (TopologySnapshot, error) { - _, services, err := disc.DiscoverAll(ctx) + services, err := disc.DiscoverAll(ctx) if err != nil { return TopologySnapshot{}, err } From 863557e20e0852b798270f757c43b330a73eec21 Mon Sep 17 00:00:00 2001 From: Benjamin Date: Wed, 30 Sep 2026 19:17:09 +0200 Subject: [PATCH 12/54] fix(alerts): keep trigger state and policy scopes on update, show toggle errors - Triggers: an absent "enabled" means enabled on creation and keeps the stored value on update, in REST and MCP alike. - Triggers: the advanced-filters gate applies to a scope filter the request changes, so a downgraded instance can still switch a scoped trigger on or off; a failed toggle or delete is shown in the trigger list. - Escalation editor: loads, shows and sends back the policy's scopes instead of wiping them on save. - AcknowledgeButton labels are in English, like the rest of the interface. --- .../components/escalation/PolicyEditor.vue | 38 +++++--- .../escalation/__tests__/PolicyEditor.spec.ts | 83 +++++++++++++++++ frontend/src/components/TriggerManager.vue | 39 +++++--- .../__tests__/TriggerManager.spec.ts | 91 +++++++++++++++++++ .../src/components/ui/AcknowledgeButton.vue | 6 +- .../ui/__tests__/AcknowledgeButton.spec.ts | 49 ++++++++++ internal/api/v1/alert_triggers.go | 15 +-- internal/api/v1/alert_triggers_test.go | 55 +++++++++++ internal/mcp/tools_triggers.go | 26 ++++-- internal/mcp/tools_triggers_test.go | 78 +++++++++++++++- 10 files changed, 429 insertions(+), 51 deletions(-) create mode 100644 frontend/src/commercial/components/escalation/__tests__/PolicyEditor.spec.ts create mode 100644 frontend/src/components/__tests__/TriggerManager.spec.ts create mode 100644 frontend/src/components/ui/__tests__/AcknowledgeButton.spec.ts diff --git a/frontend/src/commercial/components/escalation/PolicyEditor.vue b/frontend/src/commercial/components/escalation/PolicyEditor.vue index 6a178045..bea5ff38 100644 --- a/frontend/src/commercial/components/escalation/PolicyEditor.vue +++ b/frontend/src/commercial/components/escalation/PolicyEditor.vue @@ -10,7 +10,12 @@ import { RouterLink } from 'vue-router' import { useEscalationStore } from '@/commercial/stores/escalation' import { useTriggersStore } from '@/stores/triggers' import { apiFetch } from '@/services/apiFetch' -import type { EscalationPolicy, OverlapWarning as OverlapWarningType } from '@/commercial/types/escalation' +import type { + EscalationPolicy, + EscalationScope, + OverlapWarning as OverlapWarningType, + PolicyRequest, +} from '@/commercial/types/escalation' import { X, Plus, Loader2, ArrowRight, Shield } from 'lucide-vue-next' import LevelEditor from './LevelEditor.vue' import OverlapWarningComponent from './OverlapWarning.vue' @@ -46,6 +51,7 @@ const escalationApi = useEscalationApi() const name = ref(props.policy?.name ?? '') const active = ref(props.policy?.active ?? true) const severities = ref(props.policy?.filters.severities ?? []) +const scopes = ref(props.policy?.filters.scopes.map((s) => ({ ...s })) ?? []) const maxLevels = computed(() => props.maxLevels ?? 5) const levels = ref>( @@ -61,13 +67,13 @@ const overlapWarnings = ref([]) let debounceTimer: ReturnType | null = null -function buildCurrentPayload() { +function buildCurrentPayload(): PolicyRequest { return { name: name.value.trim(), active: active.value, filters: { severities: severities.value, - scopes: [], + scopes: scopes.value, }, levels: levels.value.map((l) => ({ delay_seconds: l.delay_seconds, @@ -130,18 +136,7 @@ async function handleSave() { saveError.value = null saving.value = true try { - const payload = { - name: name.value.trim(), - active: active.value, - filters: { - severities: severities.value, - scopes: [], - }, - levels: levels.value.map((l) => ({ - delay_seconds: l.delay_seconds, - channel_ids: l.channel_ids, - })), - } + const payload = buildCurrentPayload() if (props.policy) { await store.updatePolicy(props.policy.id, payload) } else { @@ -223,6 +218,19 @@ onMounted(() => { +
+ +
+ {{ s.kind }}:{{ s.ref_id }} +
+
+ ({ + useEscalationStore: () => ({ updatePolicy, createPolicy: vi.fn() }), +})) + +vi.mock('@/stores/triggers', () => ({ + useTriggersStore: () => ({ triggersForChannel: () => [], fetchTriggers: vi.fn() }), +})) + +vi.mock('@/commercial/composables/useEscalationApi', () => ({ + useEscalationApi: () => ({ overlapProbe: vi.fn().mockResolvedValue({ overlapping: [] }) }), +})) + +vi.mock('@/services/apiFetch', async (importOriginal) => ({ + ...(await importOriginal()), + apiFetch: vi.fn().mockResolvedValue({ + channels: [{ id: 'c1', name: 'ops', type: 'webhook', enabled: true }], + }), +})) + +const policy: EscalationPolicy = { + id: 'p1', + name: 'Database on-call', + active: true, + filters: { + severities: ['critical'], + scopes: [ + { kind: 'container', ref_id: 'db-1' }, + { kind: 'endpoint', ref_id: 'api-7' }, + ], + }, + levels: [{ order: 0, delay_seconds: 300, channel_ids: ['c1'] }], + created_at: '', + updated_at: '', +} + +describe('PolicyEditor', () => { + beforeEach(() => { + updatePolicy.mockReset() + updatePolicy.mockResolvedValue(policy) + }) + + it('shows the scopes of the policy it edits', async () => { + const wrapper = mount(PolicyEditor, { props: { policy }, global: { stubs: { RouterLink: RouterLinkStub } } }) + await flushPromises() + + const scopes = wrapper.find('[data-test="policy-scopes"]') + expect(scopes.exists()).toBe(true) + expect(scopes.text()).toContain('container:db-1') + expect(scopes.text()).toContain('endpoint:api-7') + }) + + it('saves the policy with its scopes', async () => { + const wrapper = mount(PolicyEditor, { props: { policy }, global: { stubs: { RouterLink: RouterLinkStub } } }) + await flushPromises() + + const save = wrapper.findAll('button').find((b) => b.text().includes('Save policy')) + expect(save).toBeDefined() + await save!.trigger('click') + await flushPromises() + + expect(updatePolicy).toHaveBeenCalledOnce() + const [id, req] = updatePolicy.mock.calls[0]! + expect(id).toBe('p1') + expect(req.filters).toEqual({ + severities: ['critical'], + scopes: [ + { kind: 'container', ref_id: 'db-1' }, + { kind: 'endpoint', ref_id: 'api-7' }, + ], + }) + }) +}) diff --git a/frontend/src/components/TriggerManager.vue b/frontend/src/components/TriggerManager.vue index 41165e30..b5b911bc 100644 --- a/frontend/src/components/TriggerManager.vue +++ b/frontend/src/components/TriggerManager.vue @@ -20,6 +20,7 @@ const confirm = useConfirm() const showEditor = ref(false) const editingTrigger = ref(null) +const actionError = ref(null) function openCreate() { editingTrigger.value = null @@ -49,19 +50,32 @@ async function handleDelete(t: AlertTrigger) { destructive: true, }) if (!ok) return - await store.remove(t.id) + await runAction(() => store.remove(t.id), 'Failed to delete trigger.') } async function handleToggleEnabled(t: AlertTrigger) { - await store.update(t.id, { - name: t.name, - filter_severities: t.filter_severities, - filter_sources: t.filter_sources, - filter_scopes: t.filter_scopes, - enabled: !t.enabled, - notify_on_resolve: t.notify_on_resolve, - channel_ids: t.channel_ids, - }) + await runAction( + () => + store.update(t.id, { + name: t.name, + filter_severities: t.filter_severities, + filter_sources: t.filter_sources, + filter_scopes: t.filter_scopes, + enabled: !t.enabled, + notify_on_resolve: t.notify_on_resolve, + channel_ids: t.channel_ids, + }), + t.enabled ? 'Failed to disable trigger.' : 'Failed to enable trigger.', + ) +} + +async function runAction(action: () => Promise, fallback: string) { + actionError.value = null + try { + await action() + } catch (e) { + actionError.value = e instanceof Error ? e.message : fallback + } } onMounted(async () => { @@ -113,10 +127,11 @@ onMounted(async () => { diff --git a/frontend/src/components/__tests__/TriggerManager.spec.ts b/frontend/src/components/__tests__/TriggerManager.spec.ts new file mode 100644 index 00000000..746f3fbf --- /dev/null +++ b/frontend/src/components/__tests__/TriggerManager.spec.ts @@ -0,0 +1,91 @@ +// Copyright 2026 Benjamin Touchard (kOlapsis) +// SPDX-License-Identifier: Apache-2.0 + +import { describe, it, expect, vi, beforeEach } from 'vitest' +import { mount, flushPromises, RouterLinkStub } from '@vue/test-utils' +import type { AlertTrigger } from '@/types/triggers' +import { ApiError } from '@/services/apiFetch' +import TriggerManager from '@/components/TriggerManager.vue' +import TriggerList from '@/components/TriggerList.vue' + +const trigger: AlertTrigger = { + id: 't1', + name: 'Scoped', + filter_severities: 'critical', + filter_sources: '', + filter_scopes: 'container:42', + enabled: true, + notify_on_resolve: false, + channel_ids: ['c1'], + created_at: '', + updated_at: '', +} + +const update = vi.fn() + +vi.mock('@/stores/triggers', () => ({ + useTriggersStore: () => ({ + triggers: [trigger], + loading: false, + error: null, + fetchTriggers: vi.fn(), + update, + remove: vi.fn(), + triggersForChannel: () => [], + }), +})) + +vi.mock('@/stores/channels', () => ({ + useChannelsStore: () => ({ channels: [], fetchChannels: vi.fn() }), +})) + +vi.mock('@/composables/useConfirm', () => ({ useConfirm: () => vi.fn() })) + +function mountManager() { + return mount(TriggerManager, { global: { stubs: { RouterLink: RouterLinkStub } } }) +} + +describe('TriggerManager', () => { + beforeEach(() => { + update.mockReset() + }) + + it('switches a trigger off without touching its filters', async () => { + update.mockResolvedValue({ ...trigger, enabled: false }) + const wrapper = mountManager() + await flushPromises() + + wrapper.findComponent(TriggerList).vm.$emit('toggle', trigger) + await flushPromises() + + expect(update).toHaveBeenCalledWith('t1', { + name: 'Scoped', + filter_severities: 'critical', + filter_sources: '', + filter_scopes: 'container:42', + enabled: false, + notify_on_resolve: false, + channel_ids: ['c1'], + }) + expect(wrapper.find('[role="alert"]').exists()).toBe(false) + }) + + it('shows a refused toggle instead of swallowing it', async () => { + update.mockRejectedValue( + new ApiError( + 403, + { code: 'EDITION_REQUIRED', message: 'This feature requires the Personal edition.' }, + 'This feature requires the Personal edition.', + ), + ) + const wrapper = mountManager() + await flushPromises() + + wrapper.findComponent(TriggerList).vm.$emit('toggle', trigger) + await flushPromises() + + const alert = wrapper.find('[role="alert"]') + expect(alert.exists()).toBe(true) + expect(alert.text()).toBe('This feature requires the Personal edition.') + }) +}) diff --git a/frontend/src/components/ui/AcknowledgeButton.vue b/frontend/src/components/ui/AcknowledgeButton.vue index 8415860a..ba528e28 100644 --- a/frontend/src/components/ui/AcknowledgeButton.vue +++ b/frontend/src/components/ui/AcknowledgeButton.vue @@ -35,12 +35,12 @@ async function acknowledge() { - Acquittée + Acknowledged @@ -54,7 +54,7 @@ async function acknowledge() { - {{ pending ? 'Acquittement…' : 'Acquitter' }} + {{ pending ? 'Acknowledging…' : 'Acknowledge' }} diff --git a/frontend/src/components/ui/__tests__/AcknowledgeButton.spec.ts b/frontend/src/components/ui/__tests__/AcknowledgeButton.spec.ts new file mode 100644 index 00000000..a8b9584e --- /dev/null +++ b/frontend/src/components/ui/__tests__/AcknowledgeButton.spec.ts @@ -0,0 +1,49 @@ +// Copyright 2026 Benjamin Touchard (kOlapsis) +// SPDX-License-Identifier: Apache-2.0 + +import { describe, it, expect, vi } from 'vitest' +import { mount, flushPromises } from '@vue/test-utils' +import type { Alert } from '@/services/alertApi' +import AcknowledgeButton from '@/components/ui/AcknowledgeButton.vue' + +let release: () => void = () => {} +const acknowledgeAlert = vi.fn(() => new Promise((resolve) => (release = resolve))) + +vi.mock('@/stores/alerts', () => ({ + useAlertsStore: () => ({ acknowledgeAlert }), +})) + +function alertWith(over: Partial): Alert { + return { id: 'a1', status: 'active', ...over } as Alert +} + +describe('AcknowledgeButton', () => { + it('offers to acknowledge, then says it is doing so', async () => { + const wrapper = mount(AcknowledgeButton, { props: { alert: alertWith({}) } }) + const button = wrapper.find('button') + expect(button.text()).toBe('Acknowledge') + + await button.trigger('click') + expect(button.text()).toBe('Acknowledging…') + + release() + await flushPromises() + expect(button.text()).toBe('Acknowledge') + }) + + it('names who acknowledged the alert', () => { + const wrapper = mount(AcknowledgeButton, { + props: { alert: alertWith({ acknowledged_at: '2026-09-30T12:00:00Z', acknowledged_by: 'alice' }) }, + }) + const done = wrapper.find('.ack-done') + expect(done.text()).toBe('Acknowledged') + expect(done.attributes('title')).toBe('Acknowledged by alice') + }) + + it('says acknowledged when nobody is named', () => { + const wrapper = mount(AcknowledgeButton, { + props: { alert: alertWith({ acknowledged_at: '2026-09-30T12:00:00Z' }) }, + }) + expect(wrapper.find('.ack-done').attributes('title')).toBe('Acknowledged') + }) +}) diff --git a/internal/api/v1/alert_triggers.go b/internal/api/v1/alert_triggers.go index c173b64b..5c370330 100644 --- a/internal/api/v1/alert_triggers.go +++ b/internal/api/v1/alert_triggers.go @@ -87,7 +87,7 @@ func (h *AlertTriggerHandler) HandleCreateTrigger(w http.ResponseWriter, r *http WriteError(w, http.StatusBadRequest, "validation_failed", err.Error()) return } - if refuseAdvancedFilters(w, &input) { + if refuseAdvancedFilters(w, input.FilterScopes, "") { return } if err := h.checkChannelsExist(r, input.ChannelIDs); err != nil { @@ -164,7 +164,7 @@ func (h *AlertTriggerHandler) HandleUpdateTrigger(w http.ResponseWriter, r *http WriteError(w, http.StatusBadRequest, "validation_failed", err.Error()) return } - if refuseAdvancedFilters(w, &input) { + if refuseAdvancedFilters(w, input.FilterScopes, existing.FilterScopes) { return } if err := h.checkChannelsExist(r, input.ChannelIDs); err != nil { @@ -172,7 +172,7 @@ func (h *AlertTriggerHandler) HandleUpdateTrigger(w http.ResponseWriter, r *http return } - enabled := true + enabled := existing.Enabled if input.Enabled != nil { enabled = *input.Enabled } @@ -251,10 +251,11 @@ func validateTriggerInput(t *triggerInput) error { return nil } -// refuseAdvancedFilters writes the edition refusal when a scope filter is set -// on an edition that does not open it, and reports whether it did. -func refuseAdvancedFilters(w http.ResponseWriter, t *triggerInput) bool { - if t.FilterScopes == "" || extension.Allows(extension.CapAlertAdvancedFilters) { +// refuseAdvancedFilters writes the edition refusal when a request sets a scope +// filter other than the stored one on an edition that does not open it, and +// reports whether it did. +func refuseAdvancedFilters(w http.ResponseWriter, scopes, stored string) bool { + if scopes == "" || scopes == stored || extension.Allows(extension.CapAlertAdvancedFilters) { return false } refuseCapability(w, extension.CapAlertAdvancedFilters) diff --git a/internal/api/v1/alert_triggers_test.go b/internal/api/v1/alert_triggers_test.go index fc22d18f..0f6fa3fa 100644 --- a/internal/api/v1/alert_triggers_test.go +++ b/internal/api/v1/alert_triggers_test.go @@ -354,6 +354,61 @@ func TestHandleUpdateTrigger_NotifyOnResolve_OmittedKeepsExisting(t *testing.T) assert.False(t, got.NotifyOnResolve, "omitted field must keep the existing value") } +func TestHandleUpdateTrigger_Enabled_OmittedKeepsExisting(t *testing.T) { + h, ts := newTriggerHandler(true) + id := seedTrigger(t, ts, "Original") + ts.triggers[id].Enabled = false + + body := `{"name":"Renamed","channel_ids":["1"]}` + req := httptest.NewRequest("PUT", "/api/v1/alert-triggers/1", strings.NewReader(body)) + req.Header.Set("Content-Type", "application/json") + req.SetPathValue("id", id) + rec := httptest.NewRecorder() + h.HandleUpdateTrigger(rec, req) + require.Equal(t, http.StatusOK, rec.Code) + assert.False(t, ts.triggers[id].Enabled, "an update that does not mention enabled must not switch the trigger on") +} + +func TestHandleUpdateTrigger_UnchangedScopes_AllowedAfterDowngrade(t *testing.T) { + original := extension.CurrentEdition + extension.CurrentEdition = func() extension.Edition { return extension.Community } + defer func() { extension.CurrentEdition = original }() + + h, ts := newTriggerHandler(true) + id := seedTrigger(t, ts, "Scoped") + ts.triggers[id].FilterScopes = "container:42" + + for _, enabled := range []bool{false, true} { + body := fmt.Sprintf(`{"name":"Scoped","filter_scopes":"container:42","enabled":%t,"channel_ids":["1"]}`, enabled) + req := httptest.NewRequest("PUT", "/api/v1/alert-triggers/1", strings.NewReader(body)) + req.Header.Set("Content-Type", "application/json") + req.SetPathValue("id", id) + rec := httptest.NewRecorder() + h.HandleUpdateTrigger(rec, req) + require.Equal(t, http.StatusOK, rec.Code, rec.Body.String()) + assert.Equal(t, enabled, ts.triggers[id].Enabled) + } +} + +func TestHandleUpdateTrigger_ChangedScopes_RefusedAfterDowngrade(t *testing.T) { + original := extension.CurrentEdition + extension.CurrentEdition = func() extension.Edition { return extension.Community } + defer func() { extension.CurrentEdition = original }() + + h, ts := newTriggerHandler(true) + id := seedTrigger(t, ts, "Scoped") + ts.triggers[id].FilterScopes = "container:42" + + body := `{"name":"Scoped","filter_scopes":"container:43","channel_ids":["1"]}` + req := httptest.NewRequest("PUT", "/api/v1/alert-triggers/1", strings.NewReader(body)) + req.Header.Set("Content-Type", "application/json") + req.SetPathValue("id", id) + rec := httptest.NewRecorder() + h.HandleUpdateTrigger(rec, req) + assertAdvancedFiltersRefusal(t, rec) + assert.Equal(t, "container:42", ts.triggers[id].FilterScopes) +} + func TestHandleUpdateTrigger_NotFound(t *testing.T) { h, _ := newTriggerHandler(true) body := `{"name":"X","channel_ids":["1"]}` diff --git a/internal/mcp/tools_triggers.go b/internal/mcp/tools_triggers.go index 5e0198c8..3a905ab7 100644 --- a/internal/mcp/tools_triggers.go +++ b/internal/mcp/tools_triggers.go @@ -65,7 +65,7 @@ type triggerInput struct { FilterSeverities string `json:"filter_severities" jsonschema:"CSV severity filter (e.g. 'critical,warning'). Empty matches everything."` FilterSources string `json:"filter_sources" jsonschema:"CSV source filter (e.g. 'container,endpoint'). Empty matches everything."` FilterScopes string `json:"filter_scopes" jsonschema:"CSV scope filter (e.g. 'container:42,endpoint:7'). Needs the advanced filters capability. Empty matches everything."` - Enabled bool `json:"enabled" jsonschema:"Whether the trigger is active"` + Enabled *bool `json:"enabled,omitempty" jsonschema:"Whether the trigger is active. Defaults to true on creation; omitted on update, the current value is kept."` NotifyOnResolve *bool `json:"notify_on_resolve,omitempty" jsonschema:"Relay recovery (resolved) notifications, default true"` ChannelIDs []string `json:"channel_ids" jsonschema:"Notification channel IDs (at least one required)"` } @@ -86,11 +86,11 @@ func scopeFiltersRequire() string { return " Scope filters require the " + titleEdition(extension.MinEdition(extension.CapAlertAdvancedFilters)) + " edition." } -// checkAdvancedFilters refuses scope filters the running edition does not -// open. Only the "are advanced filters even in play" shortcut lives here; the -// edition decision itself goes through the registry like every other. -func checkAdvancedFilters(scopes string) (*gomcp.CallToolResult, any, error) { - if scopes == "" { +// checkAdvancedFilters refuses a scope filter other than the stored one when the +// running edition does not open it. Only the "are advanced filters even in play" +// shortcut lives here; the edition decision itself goes through the registry. +func checkAdvancedFilters(scopes, stored string) (*gomcp.CallToolResult, any, error) { + if scopes == "" || scopes == stored { return nil, nil, nil } return checkCapability(extension.CapAlertAdvancedFilters) @@ -148,7 +148,7 @@ func createTriggerHandler(svc *Services) gomcp.ToolHandlerFor[triggerInput, any] if svc.Triggers == nil || svc.Channels == nil { return errResult("trigger or channel store not available") } - if r, v, err := checkAdvancedFilters(input.FilterScopes); r != nil { + if r, v, err := checkAdvancedFilters(input.FilterScopes, ""); r != nil { return r, v, err } if err := validateTriggerCommon(&input); err != nil { @@ -164,6 +164,10 @@ func createTriggerHandler(svc *Services) gomcp.ToolHandlerFor[triggerInput, any] } } + enabled := true + if input.Enabled != nil { + enabled = *input.Enabled + } notifyOnResolve := true if input.NotifyOnResolve != nil { notifyOnResolve = *input.NotifyOnResolve @@ -174,7 +178,7 @@ func createTriggerHandler(svc *Services) gomcp.ToolHandlerFor[triggerInput, any] FilterSeverities: input.FilterSeverities, FilterSources: input.FilterSources, FilterScopes: input.FilterScopes, - Enabled: input.Enabled, + Enabled: enabled, NotifyOnResolve: notifyOnResolve, ChannelIDs: input.ChannelIDs, } @@ -201,7 +205,7 @@ func updateTriggerHandler(svc *Services) gomcp.ToolHandlerFor[updateTriggerInput if existing == nil { return errResult("trigger not found") } - if r, v, err := checkAdvancedFilters(input.FilterScopes); r != nil { + if r, v, err := checkAdvancedFilters(input.FilterScopes, existing.FilterScopes); r != nil { return r, v, err } if err := validateTriggerCommon(&input.triggerInput); err != nil { @@ -221,7 +225,9 @@ func updateTriggerHandler(svc *Services) gomcp.ToolHandlerFor[updateTriggerInput existing.FilterSeverities = input.FilterSeverities existing.FilterSources = input.FilterSources existing.FilterScopes = input.FilterScopes - existing.Enabled = input.Enabled + if input.Enabled != nil { + existing.Enabled = *input.Enabled + } if input.NotifyOnResolve != nil { existing.NotifyOnResolve = *input.NotifyOnResolve } diff --git a/internal/mcp/tools_triggers_test.go b/internal/mcp/tools_triggers_test.go index 10ce80e7..10ceeb43 100644 --- a/internal/mcp/tools_triggers_test.go +++ b/internal/mcp/tools_triggers_test.go @@ -190,7 +190,6 @@ func TestCreateTriggerHandler_Happy(t *testing.T) { result, _, err := handler(context.Background(), nil, triggerInput{ Name: "AlertAll", - Enabled: true, ChannelIDs: []string{"1"}, }) require.NoError(t, err) @@ -341,7 +340,6 @@ func TestUpdateTriggerHandler_Happy(t *testing.T) { ID: id, triggerInput: triggerInput{ Name: "AfterUpdate", - Enabled: true, ChannelIDs: []string{"1"}, }, }) @@ -363,7 +361,6 @@ func TestUpdateTriggerHandler_NotifyOnResolve_False(t *testing.T) { ID: id, triggerInput: triggerInput{ Name: "AfterUpdate", - Enabled: true, NotifyOnResolve: ¬ifyOnResolve, ChannelIDs: []string{"1"}, }, @@ -388,7 +385,6 @@ func TestUpdateTriggerHandler_NotifyOnResolve_OmittedKeepsExisting(t *testing.T) ID: id, triggerInput: triggerInput{ Name: "AfterUpdate", - Enabled: true, ChannelIDs: []string{"1"}, }, }) @@ -400,6 +396,80 @@ func TestUpdateTriggerHandler_NotifyOnResolve_OmittedKeepsExisting(t *testing.T) assert.False(t, got.NotifyOnResolve, "omitted field must keep the existing value") } +func TestCreateTriggerHandler_Enabled_DefaultsTrue(t *testing.T) { + svc, ts := buildTriggerServices() + result, _, err := createTriggerHandler(svc)(context.Background(), nil, triggerInput{ + Name: "NoEnabledField", + ChannelIDs: []string{"1"}, + }) + require.NoError(t, err) + require.False(t, result.IsError) + require.NotNil(t, ts.triggers["1"]) + assert.True(t, ts.triggers["1"].Enabled) +} + +func TestUpdateTriggerHandler_Enabled(t *testing.T) { + off := false + for name, tc := range map[string]struct { + input *bool + want bool + }{ + "omitted keeps the stored value": {nil, true}, + "explicit false switches off": {&off, false}, + } { + t.Run(name, func(t *testing.T) { + svc, ts := buildTriggerServices() + id, err := ts.InsertTrigger(context.Background(), &alert.AlertTrigger{ + Name: "On", Enabled: true, ChannelIDs: []string{"1"}, + }) + require.NoError(t, err) + + result, _, err := updateTriggerHandler(svc)(context.Background(), nil, updateTriggerInputWithID{ + ID: id, + triggerInput: triggerInput{Name: "On", Enabled: tc.input, ChannelIDs: []string{"1"}}, + }) + require.NoError(t, err) + require.False(t, result.IsError) + assert.Equal(t, tc.want, ts.triggers[id].Enabled) + }) + } +} + +func TestUpdateTriggerHandler_Scopes_AfterDowngrade(t *testing.T) { + withEdition(t, extension.Community) + off := false + for name, tc := range map[string]struct { + scopes string + refused bool + }{ + "unchanged scopes are kept": {"container:42", false}, + "new scopes are refused": {"container:43", true}, + } { + t.Run(name, func(t *testing.T) { + svc, ts := buildTriggerServices() + id, err := ts.InsertTrigger(context.Background(), &alert.AlertTrigger{ + Name: "Scoped", FilterScopes: "container:42", Enabled: true, ChannelIDs: []string{"1"}, + }) + require.NoError(t, err) + + result, _, err := updateTriggerHandler(svc)(context.Background(), nil, updateTriggerInputWithID{ + ID: id, + triggerInput: triggerInput{ + Name: "Scoped", FilterScopes: tc.scopes, Enabled: &off, ChannelIDs: []string{"1"}, + }, + }) + require.NoError(t, err) + require.Equal(t, tc.refused, result.IsError, textFromContent(t, result.Content)) + if tc.refused { + assert.Contains(t, textFromContent(t, result.Content), "edition_required") + assert.True(t, ts.triggers[id].Enabled) + } else { + assert.False(t, ts.triggers[id].Enabled) + } + }) + } +} + func TestUpdateTriggerHandler_NotFound(t *testing.T) { svc, _ := buildTriggerServices() handler := updateTriggerHandler(svc) From 4798dde6808baedb406a72fb95dd502a639c5852 Mon Sep 17 00:00:00 2001 From: Benjamin Date: Wed, 30 Sep 2026 19:18:49 +0200 Subject: [PATCH 13/54] fix(kubernetes): swap the API clients atomically on reconnect The clientset and the metrics client were plain fields, rewritten by every connection while logs, topology, stats and the event stream read them from other goroutines. They now live together in an atomic.Pointer, set once per connection and read through a single accessor; before the first connection the accessor returns an error instead of a nil client. --- internal/kubernetes/discovery.go | 12 +++-- internal/kubernetes/discovery_test.go | 65 +++++++++++----------- internal/kubernetes/events.go | 9 ++-- internal/kubernetes/events_test.go | 5 +- internal/kubernetes/health.go | 12 ++++- internal/kubernetes/health_test.go | 52 +++++++++--------- internal/kubernetes/logs.go | 18 +++++-- internal/kubernetes/logs_test.go | 78 +++++++++++++-------------- internal/kubernetes/metrics.go | 22 +++++--- internal/kubernetes/namespace.go | 6 ++- internal/kubernetes/node.go | 6 ++- internal/kubernetes/runtime.go | 45 +++++++++++++--- internal/kubernetes/runtime_test.go | 42 +++++++++++++++ internal/kubernetes/stats.go | 24 ++++++--- internal/kubernetes/workload.go | 54 ++++++++++++++----- 15 files changed, 295 insertions(+), 155 deletions(-) diff --git a/internal/kubernetes/discovery.go b/internal/kubernetes/discovery.go index cae7ad3f..3e968288 100644 --- a/internal/kubernetes/discovery.go +++ b/internal/kubernetes/discovery.go @@ -18,12 +18,16 @@ import ( // discoverAll lists Deployments, StatefulSets, DaemonSets, and bare pods. func (r *Runtime) discoverAll(ctx context.Context) ([]*cmodel.Container, error) { + cs, err := r.client() + if err != nil { + return nil, err + } now := time.Now() var containers []*cmodel.Container var rbacDenied int // Deployments - depList, err := r.clientset.AppsV1().Deployments("").List(ctx, metav1.ListOptions{}) + depList, err := cs.AppsV1().Deployments("").List(ctx, metav1.ListOptions{}) if err != nil { if k8serrors.IsForbidden(err) { r.logger.Warn("RBAC: forbidden to list deployments, skipping", "error", err) @@ -42,7 +46,7 @@ func (r *Runtime) discoverAll(ctx context.Context) ([]*cmodel.Container, error) } // StatefulSets - ssList, err := r.clientset.AppsV1().StatefulSets("").List(ctx, metav1.ListOptions{}) + ssList, err := cs.AppsV1().StatefulSets("").List(ctx, metav1.ListOptions{}) if err != nil { if k8serrors.IsForbidden(err) { r.logger.Warn("RBAC: forbidden to list statefulsets, skipping", "error", err) @@ -61,7 +65,7 @@ func (r *Runtime) discoverAll(ctx context.Context) ([]*cmodel.Container, error) } // DaemonSets - dsList, err := r.clientset.AppsV1().DaemonSets("").List(ctx, metav1.ListOptions{}) + dsList, err := cs.AppsV1().DaemonSets("").List(ctx, metav1.ListOptions{}) if err != nil { if k8serrors.IsForbidden(err) { r.logger.Warn("RBAC: forbidden to list daemonsets, skipping", "error", err) @@ -80,7 +84,7 @@ func (r *Runtime) discoverAll(ctx context.Context) ([]*cmodel.Container, error) } // Bare pods (no ownerReference to a controller) - podList, err := r.clientset.CoreV1().Pods("").List(ctx, metav1.ListOptions{}) + podList, err := cs.CoreV1().Pods("").List(ctx, metav1.ListOptions{}) if err != nil { if k8serrors.IsForbidden(err) { r.logger.Warn("RBAC: forbidden to list pods, skipping", "error", err) diff --git a/internal/kubernetes/discovery_test.go b/internal/kubernetes/discovery_test.go index 9e17900c..9e3ac4df 100644 --- a/internal/kubernetes/discovery_test.go +++ b/internal/kubernetes/discovery_test.go @@ -40,13 +40,12 @@ func TestDiscoverAll_Deployments(t *testing.T) { } cs := fake.NewClientset(dep) - rt := &Runtime{ - logger: slog.Default(), - nsFilter: NewNamespaceFilter("", ""), - clientset: cs, - prevCPU: make(map[string]*cpuPrev), - stopCh: make(chan struct{}), - } + rt := withClient(&Runtime{ + logger: slog.Default(), + nsFilter: NewNamespaceFilter("", ""), + prevCPU: make(map[string]*cpuPrev), + stopCh: make(chan struct{}), + }, cs) containers, err := rt.discoverAll(context.Background()) if err != nil { @@ -105,13 +104,12 @@ func TestDiscoverAll_NamespaceFiltering(t *testing.T) { } cs := fake.NewClientset(dep) - rt := &Runtime{ - logger: slog.Default(), - nsFilter: NewNamespaceFilter("", ""), - clientset: cs, - prevCPU: make(map[string]*cpuPrev), - stopCh: make(chan struct{}), - } + rt := withClient(&Runtime{ + logger: slog.Default(), + nsFilter: NewNamespaceFilter("", ""), + prevCPU: make(map[string]*cpuPrev), + stopCh: make(chan struct{}), + }, cs) containers, err := rt.discoverAll(context.Background()) if err != nil { @@ -143,13 +141,12 @@ func TestDiscoverAll_BarePods(t *testing.T) { } cs := fake.NewClientset(pod) - rt := &Runtime{ - logger: slog.Default(), - nsFilter: NewNamespaceFilter("", ""), - clientset: cs, - prevCPU: make(map[string]*cpuPrev), - stopCh: make(chan struct{}), - } + rt := withClient(&Runtime{ + logger: slog.Default(), + nsFilter: NewNamespaceFilter("", ""), + prevCPU: make(map[string]*cpuPrev), + stopCh: make(chan struct{}), + }, cs) containers, err := rt.discoverAll(context.Background()) if err != nil { @@ -191,13 +188,12 @@ func TestDiscoverAll_ManagedPodsExcluded(t *testing.T) { } cs := fake.NewClientset(pod) - rt := &Runtime{ - logger: slog.Default(), - nsFilter: NewNamespaceFilter("", ""), - clientset: cs, - prevCPU: make(map[string]*cpuPrev), - stopCh: make(chan struct{}), - } + rt := withClient(&Runtime{ + logger: slog.Default(), + nsFilter: NewNamespaceFilter("", ""), + prevCPU: make(map[string]*cpuPrev), + stopCh: make(chan struct{}), + }, cs) containers, err := rt.discoverAll(context.Background()) if err != nil { @@ -235,13 +231,12 @@ func TestDiscoverAll_Annotations(t *testing.T) { } cs := fake.NewClientset(dep) - rt := &Runtime{ - logger: slog.Default(), - nsFilter: NewNamespaceFilter("", ""), - clientset: cs, - prevCPU: make(map[string]*cpuPrev), - stopCh: make(chan struct{}), - } + rt := withClient(&Runtime{ + logger: slog.Default(), + nsFilter: NewNamespaceFilter("", ""), + prevCPU: make(map[string]*cpuPrev), + stopCh: make(chan struct{}), + }, cs) containers, err := rt.discoverAll(context.Background()) if err != nil { diff --git a/internal/kubernetes/events.go b/internal/kubernetes/events.go index 06dc4270..19663ef4 100644 --- a/internal/kubernetes/events.go +++ b/internal/kubernetes/events.go @@ -24,9 +24,12 @@ const informerResync = 30 * time.Second func (r *Runtime) streamEvents(ctx context.Context) <-chan runtime.RuntimeEvent { out := make(chan runtime.RuntimeEvent, 128) - r.mu.Lock() - clientset := r.clientset - r.mu.Unlock() + clientset, err := r.client() + if err != nil { + r.logger.Error("kubernetes event stream not started", "error", err) + close(out) + return out + } factory := informers.NewSharedInformerFactory(clientset, informerResync) podInformer := factory.Core().V1().Pods().Informer() diff --git a/internal/kubernetes/events_test.go b/internal/kubernetes/events_test.go index 9f06648e..ef2eb624 100644 --- a/internal/kubernetes/events_test.go +++ b/internal/kubernetes/events_test.go @@ -29,14 +29,13 @@ import ( const testProbeEvery = 20 * time.Millisecond func streamRuntime(cs *fake.Clientset) *Runtime { - return &Runtime{ + return withClient(&Runtime{ logger: slog.New(slog.NewTextHandler(io.Discard, nil)), nsFilter: NewNamespaceFilter("", ""), - clientset: cs, stopCh: make(chan struct{}), probeEvery: testProbeEvery, probeMisses: 2, - } + }, cs) } func newStreamRuntime(t *testing.T, cs *fake.Clientset) *Runtime { diff --git a/internal/kubernetes/health.go b/internal/kubernetes/health.go index 742fcbed..dac55a7b 100644 --- a/internal/kubernetes/health.go +++ b/internal/kubernetes/health.go @@ -26,7 +26,11 @@ func (r *Runtime) getHealthInfo(ctx context.Context, externalID string) (*runtim } func (r *Runtime) podHealth(ctx context.Context, ns, podName string) (*runtime.HealthInfo, error) { - pod, err := r.clientset.CoreV1().Pods(ns).Get(ctx, podName, metav1.GetOptions{}) + cs, err := r.client() + if err != nil { + return nil, err + } + pod, err := cs.CoreV1().Pods(ns).Get(ctx, podName, metav1.GetOptions{}) if err != nil { return nil, fmt.Errorf("get pod %s/%s: %w", ns, podName, err) } @@ -35,12 +39,16 @@ func (r *Runtime) podHealth(ctx context.Context, ns, podName string) (*runtime.H } func (r *Runtime) controllerHealth(ctx context.Context, ns, kind, name string) (*runtime.HealthInfo, error) { + cs, err := r.client() + if err != nil { + return nil, err + } selector, err := r.controllerSelector(ctx, ns, kind, name) if err != nil { return nil, err } - podList, err := r.clientset.CoreV1().Pods(ns).List(ctx, metav1.ListOptions{ + podList, err := cs.CoreV1().Pods(ns).List(ctx, metav1.ListOptions{ LabelSelector: selector, }) if err != nil { diff --git a/internal/kubernetes/health_test.go b/internal/kubernetes/health_test.go index 92243aad..83e03c52 100644 --- a/internal/kubernetes/health_test.go +++ b/internal/kubernetes/health_test.go @@ -32,13 +32,12 @@ func TestPodHealth_Running_AllReady(t *testing.T) { } cs := fake.NewClientset(pod) - rt := &Runtime{ - logger: slog.Default(), - nsFilter: NewNamespaceFilter("", ""), - clientset: cs, - prevCPU: make(map[string]*cpuPrev), - stopCh: make(chan struct{}), - } + rt := withClient(&Runtime{ + logger: slog.Default(), + nsFilter: NewNamespaceFilter("", ""), + prevCPU: make(map[string]*cpuPrev), + stopCh: make(chan struct{}), + }, cs) hi, err := rt.podHealth(context.Background(), "default", "web") if err != nil { @@ -77,13 +76,12 @@ func TestPodHealth_ProbeFailure(t *testing.T) { } cs := fake.NewClientset(pod) - rt := &Runtime{ - logger: slog.Default(), - nsFilter: NewNamespaceFilter("", ""), - clientset: cs, - prevCPU: make(map[string]*cpuPrev), - stopCh: make(chan struct{}), - } + rt := withClient(&Runtime{ + logger: slog.Default(), + nsFilter: NewNamespaceFilter("", ""), + prevCPU: make(map[string]*cpuPrev), + stopCh: make(chan struct{}), + }, cs) hi, err := rt.podHealth(context.Background(), "default", "web") if err != nil { @@ -112,13 +110,12 @@ func TestPodHealth_NoProbe(t *testing.T) { } cs := fake.NewClientset(pod) - rt := &Runtime{ - logger: slog.Default(), - nsFilter: NewNamespaceFilter("", ""), - clientset: cs, - prevCPU: make(map[string]*cpuPrev), - stopCh: make(chan struct{}), - } + rt := withClient(&Runtime{ + logger: slog.Default(), + nsFilter: NewNamespaceFilter("", ""), + prevCPU: make(map[string]*cpuPrev), + stopCh: make(chan struct{}), + }, cs) hi, err := rt.podHealth(context.Background(), "default", "worker") if err != nil { @@ -147,13 +144,12 @@ func TestPodHealth_Pending(t *testing.T) { } cs := fake.NewClientset(pod) - rt := &Runtime{ - logger: slog.Default(), - nsFilter: NewNamespaceFilter("", ""), - clientset: cs, - prevCPU: make(map[string]*cpuPrev), - stopCh: make(chan struct{}), - } + rt := withClient(&Runtime{ + logger: slog.Default(), + nsFilter: NewNamespaceFilter("", ""), + prevCPU: make(map[string]*cpuPrev), + stopCh: make(chan struct{}), + }, cs) hi, err := rt.podHealth(context.Background(), "default", "web") if err != nil { diff --git a/internal/kubernetes/logs.go b/internal/kubernetes/logs.go index 5a018ee6..f89b62db 100644 --- a/internal/kubernetes/logs.go +++ b/internal/kubernetes/logs.go @@ -17,6 +17,10 @@ import ( // fetchLogs retrieves the last N lines of logs from a pod. // externalID format: "namespace/pod-name[/container-name]" or "namespace/Kind/name[/container-name]". func (r *Runtime) fetchLogs(ctx context.Context, externalID string, lines int, timestamps bool) ([]string, error) { + cs, err := r.client() + if err != nil { + return nil, err + } ns, podName, containerName, err := r.resolveLogTarget(ctx, externalID) if err != nil { return nil, err @@ -31,7 +35,7 @@ func (r *Runtime) fetchLogs(ctx context.Context, externalID string, lines int, t opts.Container = containerName } - stream, err := r.clientset.CoreV1().Pods(ns).GetLogs(podName, opts).Stream(ctx) + stream, err := cs.CoreV1().Pods(ns).GetLogs(podName, opts).Stream(ctx) if err != nil { return nil, fmt.Errorf("get logs %s/%s: %w", ns, podName, err) } @@ -49,6 +53,10 @@ func (r *Runtime) fetchLogs(ctx context.Context, externalID string, lines int, t // streamLogs returns a streaming reader for pod logs. func (r *Runtime) streamLogs(ctx context.Context, externalID string, lines int, timestamps bool) (io.ReadCloser, error) { + cs, err := r.client() + if err != nil { + return nil, err + } ns, podName, containerName, err := r.resolveLogTarget(ctx, externalID) if err != nil { return nil, err @@ -64,7 +72,7 @@ func (r *Runtime) streamLogs(ctx context.Context, externalID string, lines int, opts.Container = containerName } - stream, err := r.clientset.CoreV1().Pods(ns).GetLogs(podName, opts).Stream(ctx) + stream, err := cs.CoreV1().Pods(ns).GetLogs(podName, opts).Stream(ctx) if err != nil { return nil, fmt.Errorf("stream logs %s/%s: %w", ns, podName, err) } @@ -116,12 +124,16 @@ func isControllerKind(s string) bool { // findActivePod resolves a controller to one of its running pods. func (r *Runtime) findActivePod(ctx context.Context, ns, kind, name string) (string, error) { + cs, err := r.client() + if err != nil { + return "", err + } selector, err := r.controllerSelector(ctx, ns, kind, name) if err != nil { return "", err } - podList, err := r.clientset.CoreV1().Pods(ns).List(ctx, metav1.ListOptions{ + podList, err := cs.CoreV1().Pods(ns).List(ctx, metav1.ListOptions{ LabelSelector: selector, }) if err != nil { diff --git a/internal/kubernetes/logs_test.go b/internal/kubernetes/logs_test.go index 45ca2e10..aebfed9b 100644 --- a/internal/kubernetes/logs_test.go +++ b/internal/kubernetes/logs_test.go @@ -16,13 +16,12 @@ import ( func TestResolveLogTarget_PodLevel(t *testing.T) { cs := fake.NewClientset() - rt := &Runtime{ - logger: slog.Default(), - clientset: cs, - nsFilter: NewNamespaceFilter("", ""), - prevCPU: make(map[string]*cpuPrev), - stopCh: make(chan struct{}), - } + rt := withClient(&Runtime{ + logger: slog.Default(), + nsFilter: NewNamespaceFilter("", ""), + prevCPU: make(map[string]*cpuPrev), + stopCh: make(chan struct{}), + }, cs) ns, pod, container, err := rt.resolveLogTarget(context.Background(), "default/my-pod") if err != nil { @@ -41,13 +40,12 @@ func TestResolveLogTarget_PodLevel(t *testing.T) { func TestResolveLogTarget_PodWithContainer(t *testing.T) { cs := fake.NewClientset() - rt := &Runtime{ - logger: slog.Default(), - clientset: cs, - nsFilter: NewNamespaceFilter("", ""), - prevCPU: make(map[string]*cpuPrev), - stopCh: make(chan struct{}), - } + rt := withClient(&Runtime{ + logger: slog.Default(), + nsFilter: NewNamespaceFilter("", ""), + prevCPU: make(map[string]*cpuPrev), + stopCh: make(chan struct{}), + }, cs) ns, pod, container, err := rt.resolveLogTarget(context.Background(), "default/my-pod/sidecar") if err != nil { @@ -89,13 +87,12 @@ func TestResolveLogTarget_ControllerResolvesToPod(t *testing.T) { } cs := fake.NewClientset(dep, pod) - rt := &Runtime{ - logger: slog.Default(), - clientset: cs, - nsFilter: NewNamespaceFilter("", ""), - prevCPU: make(map[string]*cpuPrev), - stopCh: make(chan struct{}), - } + rt := withClient(&Runtime{ + logger: slog.Default(), + nsFilter: NewNamespaceFilter("", ""), + prevCPU: make(map[string]*cpuPrev), + stopCh: make(chan struct{}), + }, cs) ns, podName, container, err := rt.resolveLogTarget(context.Background(), "prod/Deployment/web") if err != nil { @@ -146,13 +143,12 @@ func TestResolveLogTarget_ControllerWithContainer(t *testing.T) { } cs := fake.NewClientset(dep, pod) - rt := &Runtime{ - logger: slog.Default(), - clientset: cs, - nsFilter: NewNamespaceFilter("", ""), - prevCPU: make(map[string]*cpuPrev), - stopCh: make(chan struct{}), - } + rt := withClient(&Runtime{ + logger: slog.Default(), + nsFilter: NewNamespaceFilter("", ""), + prevCPU: make(map[string]*cpuPrev), + stopCh: make(chan struct{}), + }, cs) ns, podName, container, err := rt.resolveLogTarget(context.Background(), "prod/Deployment/web/sidecar") if err != nil { @@ -165,13 +161,12 @@ func TestResolveLogTarget_ControllerWithContainer(t *testing.T) { func TestResolveLogTarget_InvalidFormat(t *testing.T) { cs := fake.NewClientset() - rt := &Runtime{ - logger: slog.Default(), - clientset: cs, - nsFilter: NewNamespaceFilter("", ""), - prevCPU: make(map[string]*cpuPrev), - stopCh: make(chan struct{}), - } + rt := withClient(&Runtime{ + logger: slog.Default(), + nsFilter: NewNamespaceFilter("", ""), + prevCPU: make(map[string]*cpuPrev), + stopCh: make(chan struct{}), + }, cs) _, _, _, err := rt.resolveLogTarget(context.Background(), "invalid") if err == nil { @@ -202,13 +197,12 @@ func TestResolveLogTarget_NoPods(t *testing.T) { } cs := fake.NewClientset(dep) - rt := &Runtime{ - logger: slog.Default(), - clientset: cs, - nsFilter: NewNamespaceFilter("", ""), - prevCPU: make(map[string]*cpuPrev), - stopCh: make(chan struct{}), - } + rt := withClient(&Runtime{ + logger: slog.Default(), + nsFilter: NewNamespaceFilter("", ""), + prevCPU: make(map[string]*cpuPrev), + stopCh: make(chan struct{}), + }, cs) _, _, _, err := rt.resolveLogTarget(context.Background(), "default/Deployment/ghost") if err == nil { diff --git a/internal/kubernetes/metrics.go b/internal/kubernetes/metrics.go index 7fd49d2b..08454dd6 100644 --- a/internal/kubernetes/metrics.go +++ b/internal/kubernetes/metrics.go @@ -43,7 +43,8 @@ type NodeResourceMetrics struct { // rate on metrics.k8s.io at O(1) per collection cycle instead of O(pods), // avoiding client-go's default 5 QPS / burst 10 throttle on large clusters. func (r *Runtime) cachedPodMetrics(ctx context.Context, namespace, name string) (*metricsv1beta1.PodMetrics, error) { - if r.metrics == nil { + mc := r.metricsClient() + if mc == nil { return nil, fmt.Errorf("metrics-server not available") } @@ -51,7 +52,7 @@ func (r *Runtime) cachedPodMetrics(ctx context.Context, namespace, name string) defer r.podMetricsMu.Unlock() if time.Since(r.podMetricsAt) > podMetricsCacheTTL || r.podMetricsCache == nil { - list, err := r.metrics.MetricsV1beta1().PodMetricses("").List(ctx, metav1.ListOptions{}) + list, err := mc.MetricsV1beta1().PodMetricses("").List(ctx, metav1.ListOptions{}) if err != nil { return nil, fmt.Errorf("list pod metrics: %w", err) } @@ -73,6 +74,10 @@ func (r *Runtime) cachedPodMetrics(ctx context.Context, namespace, name string) // GetPodMetrics queries metrics-server for a pod's CPU and memory usage. func (r *Runtime) GetPodMetrics(ctx context.Context, namespace, name string) (*PodResourceMetrics, error) { + cs, err := r.client() + if err != nil { + return nil, err + } pm, err := r.cachedPodMetrics(ctx, namespace, name) if err != nil { return nil, err @@ -86,7 +91,7 @@ func (r *Runtime) GetPodMetrics(ctx context.Context, namespace, name string) (*P // Get memory limit from pod spec. var memLimit int64 - pod, err := r.clientset.CoreV1().Pods(namespace).Get(ctx, name, metav1.GetOptions{}) + pod, err := cs.CoreV1().Pods(namespace).Get(ctx, name, metav1.GetOptions{}) if err == nil { for _, c := range pod.Spec.Containers { if lim := c.Resources.Limits.Memory(); lim != nil { @@ -107,11 +112,16 @@ func (r *Runtime) GetPodMetrics(ctx context.Context, namespace, name string) (*P // GetNodeMetrics queries metrics-server for a node's CPU and memory usage. func (r *Runtime) GetNodeMetrics(ctx context.Context, name string) (*NodeResourceMetrics, error) { - if r.metrics == nil { + cs, err := r.client() + if err != nil { + return nil, err + } + mc := r.metricsClient() + if mc == nil { return nil, fmt.Errorf("metrics-server not available") } - nm, err := r.metrics.MetricsV1beta1().NodeMetricses().Get(ctx, name, metav1.GetOptions{}) + nm, err := mc.MetricsV1beta1().NodeMetricses().Get(ctx, name, metav1.GetOptions{}) if err != nil { return nil, fmt.Errorf("get node metrics %s: %w", name, err) } @@ -121,7 +131,7 @@ func (r *Runtime) GetNodeMetrics(ctx context.Context, name string) (*NodeResourc // Get capacity from node spec. var cpuCapacity, memCapacity int64 - node, err := r.clientset.CoreV1().Nodes().Get(ctx, name, metav1.GetOptions{}) + node, err := cs.CoreV1().Nodes().Get(ctx, name, metav1.GetOptions{}) if err == nil { if cpu := node.Status.Capacity.Cpu(); cpu != nil { cpuCapacity = cpu.MilliValue() diff --git a/internal/kubernetes/namespace.go b/internal/kubernetes/namespace.go index 48b0bcb3..7c796db7 100644 --- a/internal/kubernetes/namespace.go +++ b/internal/kubernetes/namespace.go @@ -72,7 +72,11 @@ func (f *NamespaceFilter) IsAllowed(namespace string) bool { // ListNamespaces returns the allowed namespace names from the cluster. // The result respects the allowlist/blocklist configured via env vars. func (r *Runtime) ListNamespaces(ctx context.Context) ([]string, error) { - nsList, err := r.clientset.CoreV1().Namespaces().List(ctx, metav1.ListOptions{}) + cs, err := r.client() + if err != nil { + return nil, err + } + nsList, err := cs.CoreV1().Namespaces().List(ctx, metav1.ListOptions{}) if err != nil { return nil, fmt.Errorf("list namespaces: %w", err) } diff --git a/internal/kubernetes/node.go b/internal/kubernetes/node.go index caaeca1e..60a87b32 100644 --- a/internal/kubernetes/node.go +++ b/internal/kubernetes/node.go @@ -39,7 +39,11 @@ type K8sResourceQuantity struct { // ListNodes returns all cluster nodes with their resource capacity and status. func (r *Runtime) ListNodes(ctx context.Context) ([]K8sNode, error) { - nodeList, err := r.clientset.CoreV1().Nodes().List(ctx, metav1.ListOptions{}) + cs, err := r.client() + if err != nil { + return nil, err + } + nodeList, err := cs.CoreV1().Nodes().List(ctx, metav1.ListOptions{}) if err != nil { return nil, fmt.Errorf("list nodes: %w", err) } diff --git a/internal/kubernetes/runtime.go b/internal/kubernetes/runtime.go index 7209c025..06d5fd78 100644 --- a/internal/kubernetes/runtime.go +++ b/internal/kubernetes/runtime.go @@ -5,11 +5,13 @@ package kubernetes import ( "context" + "errors" "fmt" "io" "log/slog" "os" "sync" + "sync/atomic" "time" cmodel "github.com/kolapsis/maintenant/internal/container" @@ -34,11 +36,10 @@ func init() { // Runtime implements runtime.Runtime for Kubernetes. type Runtime struct { - logger *slog.Logger - nsFilter *NamespaceFilter - clientset k8s.Interface - metrics metricsv.Interface - stopCh chan struct{} + logger *slog.Logger + nsFilter *NamespaceFilter + conn atomic.Pointer[clients] + stopCh chan struct{} probeEvery time.Duration probeMisses int @@ -53,6 +54,31 @@ type Runtime struct { podMetricsAt time.Time } +// clients are the API clients of one connection, replaced as a whole when the runtime reconnects. +type clients struct { + core k8s.Interface + metrics metricsv.Interface +} + +var errNotConnected = errors.New("kubernetes runtime has never connected") + +func (r *Runtime) client() (k8s.Interface, error) { + c := r.conn.Load() + if c == nil { + return nil, errNotConnected + } + return c.core, nil +} + +// metricsClient returns the metrics-server client of the current connection, nil when there is none. +func (r *Runtime) metricsClient() metricsv.Interface { + c := r.conn.Load() + if c == nil { + return nil + } + return c.metrics +} + type cpuPrev struct { milliCPU int64 timestamp time.Time @@ -124,9 +150,8 @@ func (r *Runtime) connect(ctx context.Context, config *rest.Config) error { } } + r.conn.Store(&clients{core: clientset, metrics: metricsClient}) r.mu.Lock() - r.clientset = clientset - r.metrics = metricsClient r.metricsAvailable = metricsOK r.connected = true r.mu.Unlock() @@ -234,12 +259,16 @@ func (r *Runtime) GetHealthInfo(ctx context.Context, externalID string) (*runtim // ListContainerNames returns the container names in a workload's pod spec. // For controllers, resolves to a pod's spec. For bare pods, reads the pod directly. func (r *Runtime) ListContainerNames(ctx context.Context, externalID string) ([]string, error) { + cs, err := r.client() + if err != nil { + return nil, err + } ns, podName, _, err := r.resolveLogTarget(ctx, externalID) if err != nil { return nil, err } - pod, err := r.clientset.CoreV1().Pods(ns).Get(ctx, podName, metav1.GetOptions{}) + pod, err := cs.CoreV1().Pods(ns).Get(ctx, podName, metav1.GetOptions{}) if err != nil { return nil, fmt.Errorf("get pod %s/%s: %w", ns, podName, err) } diff --git a/internal/kubernetes/runtime_test.go b/internal/kubernetes/runtime_test.go index 8ccf4621..844099f7 100644 --- a/internal/kubernetes/runtime_test.go +++ b/internal/kubernetes/runtime_test.go @@ -11,12 +11,14 @@ import ( "net/http/httptest" "os" "path/filepath" + "sync" "sync/atomic" "testing" "time" "github.com/stretchr/testify/assert" "github.com/stretchr/testify/require" + k8s "k8s.io/client-go/kubernetes" ) // fakeAPIServer answers /version with 503 for the first failures calls, then with a version. @@ -48,6 +50,11 @@ func fakeAPIServer(t *testing.T, failures int32) *atomic.Int32 { return &hits } +func withClient(r *Runtime, cs k8s.Interface) *Runtime { + r.conn.Store(&clients{core: cs}) + return r +} + func newTestRuntime(t *testing.T) *Runtime { t.Helper() r, err := NewRuntime(slog.New(slog.NewTextHandler(io.Discard, nil)), NewNamespaceFilter("", "")) @@ -79,6 +86,41 @@ func TestConnect_StopsWithItsContext(t *testing.T) { assert.False(t, r.IsConnected()) } +// Run under -race: logs, topology and stats read the clients while reconnections replace them. +func TestReconnect_WhileReading(t *testing.T) { + fakeAPIServer(t, 0) + r := newTestRuntime(t) + ctx := context.Background() + require.NoError(t, r.TryConnect(ctx)) + + readers := []func(){ + func() { _, _ = r.FetchLogs(ctx, "default/web", 10, false) }, + func() { _, _ = SnapshotFromRuntime(ctx, r) }, + func() { _, _ = r.StatsSnapshot(ctx, "default/web") }, + } + done := make(chan struct{}) + var wg sync.WaitGroup + for _, read := range readers { + wg.Add(1) + go func() { + defer wg.Done() + for { + select { + case <-done: + return + default: + read() + } + } + }() + } + for range 20 { + require.NoError(t, r.TryConnect(ctx)) + } + close(done) + wg.Wait() +} + func TestTryConnect_MakesASingleAttempt(t *testing.T) { hits := fakeAPIServer(t, 1) r := newTestRuntime(t) diff --git a/internal/kubernetes/stats.go b/internal/kubernetes/stats.go index c0b321be..e42f43f4 100644 --- a/internal/kubernetes/stats.go +++ b/internal/kubernetes/stats.go @@ -17,7 +17,7 @@ import ( // statsSnapshot queries metrics-server for a workload's CPU and memory. // externalID format: "namespace/ControllerKind/name" or "namespace/pod-name". func (r *Runtime) statsSnapshot(ctx context.Context, externalID string) (*runtime.RawStats, error) { - if r.metrics == nil { + if r.metricsClient() == nil { return nil, fmt.Errorf("metrics-server not available") } @@ -36,6 +36,10 @@ func (r *Runtime) statsSnapshot(ctx context.Context, externalID string) (*runtim } func (r *Runtime) podStats(ctx context.Context, ns, podName string) (*runtime.RawStats, error) { + cs, err := r.client() + if err != nil { + return nil, err + } pm, err := r.cachedPodMetrics(ctx, ns, podName) if err != nil { return nil, err @@ -51,7 +55,7 @@ func (r *Runtime) podStats(ctx context.Context, ns, podName string) (*runtime.Ra cpuPercent := r.computeCPUPercent(pm.Name, totalCPUMilli, pm.Timestamp.Time) // Get memory limit from pod spec. - pod, err := r.clientset.CoreV1().Pods(ns).Get(ctx, podName, metav1.GetOptions{}) + pod, err := cs.CoreV1().Pods(ns).Get(ctx, podName, metav1.GetOptions{}) var memLimit int64 if err == nil { for _, c := range pod.Spec.Containers { @@ -74,13 +78,17 @@ func (r *Runtime) podStats(ctx context.Context, ns, podName string) (*runtime.Ra } func (r *Runtime) controllerStats(ctx context.Context, ns, kind, name string) (*runtime.RawStats, error) { + cs, err := r.client() + if err != nil { + return nil, err + } // Build label selector from controller spec. selector, err := r.controllerSelector(ctx, ns, kind, name) if err != nil { return nil, err } - podList, err := r.clientset.CoreV1().Pods(ns).List(ctx, metav1.ListOptions{ + podList, err := cs.CoreV1().Pods(ns).List(ctx, metav1.ListOptions{ LabelSelector: selector, }) if err != nil { @@ -142,9 +150,13 @@ func (r *Runtime) computeCPUPercent(key string, milliCPU int64, ts time.Time) fl } func (r *Runtime) controllerSelector(ctx context.Context, ns, kind, name string) (string, error) { + cs, err := r.client() + if err != nil { + return "", err + } switch kind { case "Deployment": - dep, err := r.clientset.AppsV1().Deployments(ns).Get(ctx, name, metav1.GetOptions{}) + dep, err := cs.AppsV1().Deployments(ns).Get(ctx, name, metav1.GetOptions{}) if err != nil { return "", fmt.Errorf("get deployment %s/%s: %w", ns, name, err) } @@ -152,7 +164,7 @@ func (r *Runtime) controllerSelector(ctx context.Context, ns, kind, name string) return labels.Set(dep.Spec.Selector.MatchLabels).String(), nil } case "StatefulSet": - ss, err := r.clientset.AppsV1().StatefulSets(ns).Get(ctx, name, metav1.GetOptions{}) + ss, err := cs.AppsV1().StatefulSets(ns).Get(ctx, name, metav1.GetOptions{}) if err != nil { return "", fmt.Errorf("get statefulset %s/%s: %w", ns, name, err) } @@ -160,7 +172,7 @@ func (r *Runtime) controllerSelector(ctx context.Context, ns, kind, name string) return labels.Set(ss.Spec.Selector.MatchLabels).String(), nil } case "DaemonSet": - ds, err := r.clientset.AppsV1().DaemonSets(ns).Get(ctx, name, metav1.GetOptions{}) + ds, err := cs.AppsV1().DaemonSets(ns).Get(ctx, name, metav1.GetOptions{}) if err != nil { return "", fmt.Errorf("get daemonset %s/%s: %w", ns, name, err) } diff --git a/internal/kubernetes/workload.go b/internal/kubernetes/workload.go index d81ac9b9..7d476948 100644 --- a/internal/kubernetes/workload.go +++ b/internal/kubernetes/workload.go @@ -107,6 +107,10 @@ type PodFilters struct { // non-empty only those namespaces are queried; otherwise all allowed // namespaces are included. func (r *Runtime) ListWorkloads(ctx context.Context, namespaces []string) ([]K8sWorkloadGroup, error) { + cs, err := r.client() + if err != nil { + return nil, err + } targetNS := r.resolveNamespaces(namespaces) // Collect per-namespace workloads. @@ -116,7 +120,7 @@ func (r *Runtime) ListWorkloads(ctx context.Context, namespaces []string) ([]K8s } // Deployments. - depList, err := r.clientset.AppsV1().Deployments("").List(ctx, metav1.ListOptions{}) + depList, err := cs.AppsV1().Deployments("").List(ctx, metav1.ListOptions{}) if err != nil { if !k8serrors.IsForbidden(err) { return nil, fmt.Errorf("list deployments: %w", err) @@ -136,7 +140,7 @@ func (r *Runtime) ListWorkloads(ctx context.Context, namespaces []string) ([]K8s } // StatefulSets. - ssList, err := r.clientset.AppsV1().StatefulSets("").List(ctx, metav1.ListOptions{}) + ssList, err := cs.AppsV1().StatefulSets("").List(ctx, metav1.ListOptions{}) if err != nil { if !k8serrors.IsForbidden(err) { return nil, fmt.Errorf("list statefulsets: %w", err) @@ -156,7 +160,7 @@ func (r *Runtime) ListWorkloads(ctx context.Context, namespaces []string) ([]K8s } // DaemonSets. - dsList, err := r.clientset.AppsV1().DaemonSets("").List(ctx, metav1.ListOptions{}) + dsList, err := cs.AppsV1().DaemonSets("").List(ctx, metav1.ListOptions{}) if err != nil { if !k8serrors.IsForbidden(err) { return nil, fmt.Errorf("list daemonsets: %w", err) @@ -176,7 +180,7 @@ func (r *Runtime) ListWorkloads(ctx context.Context, namespaces []string) ([]K8s } // Jobs. - jobList, err := r.clientset.BatchV1().Jobs("").List(ctx, metav1.ListOptions{}) + jobList, err := cs.BatchV1().Jobs("").List(ctx, metav1.ListOptions{}) if err != nil { if !k8serrors.IsForbidden(err) { return nil, fmt.Errorf("list jobs: %w", err) @@ -260,6 +264,10 @@ func (r *Runtime) GetWorkload(ctx context.Context, id string) (*K8sWorkload, []K // ListPods returns a flat pod list optionally filtered by workload, node, and status. func (r *Runtime) ListPods(ctx context.Context, namespaces []string, filters PodFilters) ([]K8sPod, error) { + cs, err := r.client() + if err != nil { + return nil, err + } targetNS := r.resolveNamespaces(namespaces) listNS := "" @@ -267,7 +275,7 @@ func (r *Runtime) ListPods(ctx context.Context, namespaces []string, filters Pod listNS = targetNS[0] } - podList, err := r.clientset.CoreV1().Pods(listNS).List(ctx, metav1.ListOptions{}) + podList, err := cs.CoreV1().Pods(listNS).List(ctx, metav1.ListOptions{}) if err != nil { return nil, fmt.Errorf("list pods: %w", err) } @@ -300,7 +308,11 @@ func (r *Runtime) ListPods(ctx context.Context, namespaces []string, filters Pod // GetPodDetail returns details for a single pod plus recent events. func (r *Runtime) GetPodDetail(ctx context.Context, namespace, name string) (*K8sPod, []K8sEvent, error) { - pod, err := r.clientset.CoreV1().Pods(namespace).Get(ctx, name, metav1.GetOptions{}) + cs, err := r.client() + if err != nil { + return nil, nil, err + } + pod, err := cs.CoreV1().Pods(namespace).Get(ctx, name, metav1.GetOptions{}) if err != nil { return nil, nil, fmt.Errorf("get pod %s/%s: %w", namespace, name, err) } @@ -319,30 +331,34 @@ func (r *Runtime) GetPodDetail(ctx context.Context, namespace, name string) (*K8 // --- internal helpers --- func (r *Runtime) fetchWorkload(ctx context.Context, ns, kind, name string) (*K8sWorkload, error) { + cs, err := r.client() + if err != nil { + return nil, err + } switch kind { case "Deployment": - dep, err := r.clientset.AppsV1().Deployments(ns).Get(ctx, name, metav1.GetOptions{}) + dep, err := cs.AppsV1().Deployments(ns).Get(ctx, name, metav1.GetOptions{}) if err != nil { return nil, fmt.Errorf("get deployment %s/%s: %w", ns, name, err) } wl := mapDeploymentWorkload(dep) return &wl, nil case "StatefulSet": - ss, err := r.clientset.AppsV1().StatefulSets(ns).Get(ctx, name, metav1.GetOptions{}) + ss, err := cs.AppsV1().StatefulSets(ns).Get(ctx, name, metav1.GetOptions{}) if err != nil { return nil, fmt.Errorf("get statefulset %s/%s: %w", ns, name, err) } wl := mapStatefulSetWorkload(ss) return &wl, nil case "DaemonSet": - ds, err := r.clientset.AppsV1().DaemonSets(ns).Get(ctx, name, metav1.GetOptions{}) + ds, err := cs.AppsV1().DaemonSets(ns).Get(ctx, name, metav1.GetOptions{}) if err != nil { return nil, fmt.Errorf("get daemonset %s/%s: %w", ns, name, err) } wl := mapDaemonSetWorkload(ds) return &wl, nil case "Job": - job, err := r.clientset.BatchV1().Jobs(ns).Get(ctx, name, metav1.GetOptions{}) + job, err := cs.BatchV1().Jobs(ns).Get(ctx, name, metav1.GetOptions{}) if err != nil { return nil, fmt.Errorf("get job %s/%s: %w", ns, name, err) } @@ -354,7 +370,11 @@ func (r *Runtime) fetchWorkload(ctx context.Context, ns, kind, name string) (*K8 } func (r *Runtime) listPodsForSelector(ctx context.Context, ns, selector, workloadID string) ([]K8sPod, error) { - podList, err := r.clientset.CoreV1().Pods(ns).List(ctx, metav1.ListOptions{ + cs, err := r.client() + if err != nil { + return nil, err + } + podList, err := cs.CoreV1().Pods(ns).List(ctx, metav1.ListOptions{ LabelSelector: selector, }) if err != nil { @@ -384,7 +404,11 @@ func (r *Runtime) listPodEvents(ctx context.Context, ns, name string) ([]K8sEven // with the object it concerns, so the server can persist them per-agent and // serve them back on workload/pod detail views. func (r *Runtime) ListAllEvents(ctx context.Context) ([]K8sEventRef, error) { - evtList, err := r.clientset.CoreV1().Events("").List(ctx, metav1.ListOptions{}) + cs, err := r.client() + if err != nil { + return nil, err + } + evtList, err := cs.CoreV1().Events("").List(ctx, metav1.ListOptions{}) if err != nil { return nil, fmt.Errorf("list all events: %w", err) } @@ -418,7 +442,11 @@ func (r *Runtime) ListAllEvents(ctx context.Context) ([]K8sEventRef, error) { } func (r *Runtime) fetchEvents(ctx context.Context, ns, fieldSelector string) ([]K8sEvent, error) { - evtList, err := r.clientset.CoreV1().Events(ns).List(ctx, metav1.ListOptions{ + cs, err := r.client() + if err != nil { + return nil, err + } + evtList, err := cs.CoreV1().Events(ns).List(ctx, metav1.ListOptions{ FieldSelector: fieldSelector, }) if err != nil { From f785c23284e38cf2a376253b448cac6c7bc45260 Mon Sep 17 00:00:00 2001 From: Benjamin Date: Wed, 30 Sep 2026 19:22:17 +0200 Subject: [PATCH 14/54] fix(kubernetes): list Services through the connection's client accessor --- internal/kubernetes/services.go | 19 ++++++++++++------- internal/kubernetes/services_test.go | 13 ++++++------- 2 files changed, 18 insertions(+), 14 deletions(-) diff --git a/internal/kubernetes/services.go b/internal/kubernetes/services.go index caf371a9..76656eea 100644 --- a/internal/kubernetes/services.go +++ b/internal/kubernetes/services.go @@ -13,6 +13,7 @@ import ( metav1 "k8s.io/apimachinery/pkg/apis/meta/v1" "k8s.io/apimachinery/pkg/labels" "k8s.io/apimachinery/pkg/util/intstr" + k8s "k8s.io/client-go/kubernetes" ) // ServiceExposure is one port a LoadBalancer or NodePort Service opens outside @@ -36,7 +37,11 @@ type selectableWorkload struct { // ListServiceExposures returns the ports LoadBalancer and NodePort Services // expose, one entry per workload each Service selects. func (r *Runtime) ListServiceExposures(ctx context.Context) ([]ServiceExposure, error) { - list, err := r.clientset.CoreV1().Services("").List(ctx, metav1.ListOptions{}) + cs, err := r.client() + if err != nil { + return nil, err + } + list, err := cs.CoreV1().Services("").List(ctx, metav1.ListOptions{}) if err != nil { return nil, fmt.Errorf("list services: %w", err) } @@ -54,7 +59,7 @@ func (r *Runtime) ListServiceExposures(ctx context.Context) ([]ServiceExposure, return nil, nil } - workloads, err := r.selectableWorkloads(ctx) + workloads, err := r.selectableWorkloads(ctx, cs) if err != nil { return nil, err } @@ -85,7 +90,7 @@ func (r *Runtime) ListServiceExposures(ctx context.Context) ([]ServiceExposure, // selectableWorkloads lists what a Service selector can match: the pod // templates of the controllers and the labels of bare pods. A kind the RBAC // denies is left out, as in discoverAll. -func (r *Runtime) selectableWorkloads(ctx context.Context) ([]selectableWorkload, error) { +func (r *Runtime) selectableWorkloads(ctx context.Context, cs k8s.Interface) ([]selectableWorkload, error) { var out []selectableWorkload keep := func(ns string) bool { return r.nsFilter.IsAllowed(ns) } skip := func(kind string, err error) error { @@ -96,7 +101,7 @@ func (r *Runtime) selectableWorkloads(ctx context.Context) ([]selectableWorkload return fmt.Errorf("list %s: %w", kind, err) } - deployments, err := r.clientset.AppsV1().Deployments("").List(ctx, metav1.ListOptions{}) + deployments, err := cs.AppsV1().Deployments("").List(ctx, metav1.ListOptions{}) if err != nil { if err := skip("deployments", err); err != nil { return nil, err @@ -109,7 +114,7 @@ func (r *Runtime) selectableWorkloads(ctx context.Context) ([]selectableWorkload } } - statefulSets, err := r.clientset.AppsV1().StatefulSets("").List(ctx, metav1.ListOptions{}) + statefulSets, err := cs.AppsV1().StatefulSets("").List(ctx, metav1.ListOptions{}) if err != nil { if err := skip("statefulsets", err); err != nil { return nil, err @@ -122,7 +127,7 @@ func (r *Runtime) selectableWorkloads(ctx context.Context) ([]selectableWorkload } } - daemonSets, err := r.clientset.AppsV1().DaemonSets("").List(ctx, metav1.ListOptions{}) + daemonSets, err := cs.AppsV1().DaemonSets("").List(ctx, metav1.ListOptions{}) if err != nil { if err := skip("daemonsets", err); err != nil { return nil, err @@ -135,7 +140,7 @@ func (r *Runtime) selectableWorkloads(ctx context.Context) ([]selectableWorkload } } - pods, err := r.clientset.CoreV1().Pods("").List(ctx, metav1.ListOptions{}) + pods, err := cs.CoreV1().Pods("").List(ctx, metav1.ListOptions{}) if err != nil { if err := skip("pods", err); err != nil { return nil, err diff --git a/internal/kubernetes/services_test.go b/internal/kubernetes/services_test.go index 6fbacf43..3000e5fc 100644 --- a/internal/kubernetes/services_test.go +++ b/internal/kubernetes/services_test.go @@ -19,13 +19,12 @@ import ( ) func serviceRuntime(nsFilter *NamespaceFilter, objects ...k8sruntime.Object) *Runtime { - return &Runtime{ - logger: slog.Default(), - nsFilter: nsFilter, - clientset: fake.NewClientset(objects...), - prevCPU: make(map[string]*cpuPrev), - stopCh: make(chan struct{}), - } + return withClient(&Runtime{ + logger: slog.Default(), + nsFilter: nsFilter, + prevCPU: make(map[string]*cpuPrev), + stopCh: make(chan struct{}), + }, fake.NewClientset(objects...)) } func service(ns, name string, typ corev1.ServiceType, selector map[string]string, ports ...corev1.ServicePort) *corev1.Service { From fc8ca37824133c1d0d7104e247500f04d41228e6 Mon Sep 17 00:00:00 2001 From: Benjamin Date: Wed, 30 Sep 2026 19:54:01 +0200 Subject: [PATCH 15/54] docs: align status page, MCP, escalation and security pages with the code --- docs/features/alert-escalation.md | 98 ++++++--- docs/features/mcp.md | 159 +++++++++------ docs/features/security.md | 105 ++++++---- docs/features/status-page.md | 327 +++++++++++++++++++++++++----- 4 files changed, 508 insertions(+), 181 deletions(-) diff --git a/docs/features/alert-escalation.md b/docs/features/alert-escalation.md index ec0550fe..17621c42 100644 --- a/docs/features/alert-escalation.md +++ b/docs/features/alert-escalation.md @@ -4,24 +4,30 @@ Escalation policies automatically route unacknowledged alerts through a chain of notification levels, each with a configurable delay and a distinct set of channels. -> **Pattern: reserved-escalation channel.** Since channels are silent by default (they only receive alerts when wired through an [Alert Trigger](alerts.md#alert-triggers)), you can create a channel that exists *only* for an escalation level. The directrice technique's email referenced in Level 3 of a policy, with no trigger using it, will *only* be notified after T+1h of unacknowledged escalation — never at the initial dispatch. This is the cleanest way to model "last-resort" destinations without duplicate notifications. +> **Pattern: reserved-escalation channel.** Since channels are silent by default (they only receive alerts when wired through an [Alert Trigger](alerts.md#alert-triggers)), you can create a channel that exists *only* for an escalation level. A manager's email referenced in Level 3 of a policy, with no trigger using it, will *only* be notified after 1 hour of unacknowledged escalation, never at the initial dispatch. This is the cleanest way to model "last-resort" destinations without duplicate notifications. Such a channel gets the escalation, acknowledgment and exhaustion notices, but no notice when the alert resolves. --- ## Concept -When an alert fires and remains unacknowledged, an active escalation policy that matches the alert will start an **escalation run**. The run tracks which notification level is currently due and dispatches notifications to the configured channels at each level's delay. The chain stops as soon as the alert is acknowledged or resolved. +When an alert fires, every active escalation policy whose filters match it starts an **escalation run**. The run tracks which notification level is due next and sends the notifications of that level when its delay has passed. The chain stops as soon as the alert is acknowledged or resolved. + +A level's `delay_seconds` is counted **from the start of the run** (the moment the alert fires), not from the previous level. A level with a delay of 900 fires 15 minutes after the alert, whatever the level before it had. Runs are evaluated once a minute, so a level can fire up to a minute after its delay. ``` -Alert fires +Alert fires (T+0) │ - ├─ Level 1 (delay 5 min) → Notify #slack-oncall + ├─ Level 1 (delay 300) → T+5 min, if still unacknowledged → Notify #slack-oncall │ - ├─ Level 2 (delay 15 min, if still unacknowledged) → Page +33-6-XX + ├─ Level 2 (delay 900) → T+15 min, if still unacknowledged → Page +33-6-XX │ - └─ Level 3 (delay 1 hour, if still unacknowledged) → Email management + └─ Level 3 (delay 3600) → T+1 hour, if still unacknowledged → Email management ``` +A run keeps the policy as it was when the run started. Editing or deactivating a policy changes the alerts raised afterwards, not the runs already in progress; deleting the policy stops them. + +Policies only see alerts that are active. An alert raised while a matching [silence rule](alerts.md) or maintenance window is in force is silenced and starts no run. If an alert becomes more severe, policies that match its new severity start a run too; runs already started continue untouched. + --- ## Examples @@ -32,18 +38,18 @@ Notify the on-call Slack channel 5 minutes after an alert fires. ```json { - "name": "Critical pager — Level 1 only", + "name": "Critical pager: Level 1 only", "active": true, "filters": { "severities": ["critical"] }, "levels": [ - { "delay_seconds": 300, "channel_ids": [1] } + { "delay_seconds": 300, "channel_ids": ["0195f3c8-5e2a-7b14-9a30-7d1c4f2e8a66"] } ] } ``` ### Multi-level policy -Escalate progressively over an hour. +Escalate progressively over an hour: the levels fire 5 minutes, 15 minutes and 1 hour after the alert. ```json { @@ -51,85 +57,115 @@ Escalate progressively over an hour. "active": true, "filters": { "severities": ["critical", "warning"] }, "levels": [ - { "delay_seconds": 300, "channel_ids": [1] }, - { "delay_seconds": 900, "channel_ids": [2] }, - { "delay_seconds": 3600, "channel_ids": [3] } + { "delay_seconds": 300, "channel_ids": ["0195f3c8-5e2a-7b14-9a30-7d1c4f2e8a66"] }, + { "delay_seconds": 900, "channel_ids": ["0195f3c9-0b71-7e52-8d04-3a9f6c1b2d40"] }, + { "delay_seconds": 3600, "channel_ids": ["0195f3c9-44d8-7a3c-b6e1-92f0d5a7c318"] } ] } ``` -### Filter by tag +### Filter by entity -Only escalate alerts tagged `prod`. +Only escalate alerts about one endpoint. A scope names an alert's entity: `kind` is the entity type (`container`, `endpoint`, `heartbeat`, `certificate`) and `ref_id` is the UUID of that entity. ```json { - "name": "Prod-only chain", + "name": "Checkout endpoint chain", "active": true, "filters": { "severities": ["critical"], - "tags": ["prod"] + "scopes": [{ "kind": "endpoint", "ref_id": "0195f3c4-11aa-7f00-8c55-2b9d0e3a4c18" }] }, "levels": [ - { "delay_seconds": 300, "channel_ids": [1] } + { "delay_seconds": 300, "channel_ids": ["0195f3c8-5e2a-7b14-9a30-7d1c4f2e8a66"] } ] } ``` +An empty `severities` or `scopes` matches everything. An alert must satisfy both filters to match. The policy editor of the dashboard edits the name, the activation, the severities (`warning`, `critical`) and the levels; it shows the scopes of a policy read-only and keeps them when you save. Scopes are set through the API. + +--- + +## Validation + +| Field | Rule | +|-------|------| +| `name` | Required, at most 120 characters | +| `levels` | At least one | +| `levels[].delay_seconds` | Between 60 and 86400 | +| Consecutive levels | Each level is at least 60 seconds after the previous one, so delays must increase | +| `levels[].channel_ids` | At least one channel ID (a UUID) per level | + +A request that breaks a rule returns `400 validation_failed` with the field in the message. The channel IDs are not checked against existing channels when the policy is saved: a level that points at a deleted channel skips it when it fires. + --- ## Interactions ### Acknowledgment stops escalation -When an alert is acknowledged, any active escalation run is immediately stopped with status `stopped_by_ack`. A notification is sent on all channels that have already received an escalation notification, informing them the alert has been acknowledged. +When an alert is acknowledged, any active escalation run is immediately stopped with status `stopped_by_ack`. Every channel that received a notification of the run (a delivery that is `sent` or still `pending`) gets one acknowledgment notice, once, even if it appeared in several levels. The notice reuses the recovery format of the channel, and its message is: + +``` +Acknowledged by at